Repository navigation
[ENH] Implement a unified hyperparameter tuning interface decoupled from model classes for v2 #2335
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: main
Are you sure you want to change the base?
Changes from 3 commits
bb7d41c
0c691db
772fa91
8301098
5133352
bac8e87
8a8f97b
4dcd61c
cdd7f85
39d5122
928b8a4
b6bf90a
4ff83c2
f126396
5788e74
761d72e
bed2247
3c8d8d1
786ac9a
52b28ab
1470d12
c341a12
ee7f848
d02820a
7d03d39
c8f13d5
a2a2a7e
2eaa29c
87f018e
f7662c3
300ec2b
f86629c
51e328c
51fe98f
78fd85c
f838efd
076338c
7a10b01
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,55 @@ | ||
| """ | ||
| Global Hyperparameter Registry. | ||
|
|
||
| This is the "common ground" — standard ranges for parameters that appear | ||
| across multiple models. When BaseModel inspects a subclass's __init__ | ||
| and finds 'hidden_size', it looks up the range HERE. | ||
|
|
||
| HOW TO EXTEND: If a new model introduces a new common parameter, | ||
| just add one line here. All models using that param name are instantly tuneable. | ||
| """ | ||
|
|
||
| from pytorch_forecasting.tuning.search_range import SearchRange | ||
|
|
||
| UNIVERSAL_PARAMS = { | ||
| "optimizer": SearchRange( | ||
| param_type="categorical", | ||
| choices=["adam", "adamw"], | ||
| ), | ||
| "optimizer_params.lr": SearchRange( | ||
| param_type="float", | ||
| low=1e-5, | ||
| high=1e-1, | ||
| log=True, | ||
| ), | ||
| } | ||
|
|
||
| MODEL_PARAMS = { | ||
| "hidden_size": SearchRange(param_type="int", low=16, high=512, log=True), | ||
| "dropout": SearchRange(param_type="float", low=0.05, high=0.5), | ||
| "dropout_rate": SearchRange(param_type="float", low=0.05, high=0.5), | ||
| "activation": SearchRange(param_type="categorical", choices=["relu", "gelu"]), | ||
| "n_heads": SearchRange(param_type="categorical", choices=[1, 2, 4, 8]), | ||
| "attention_head_size": SearchRange(param_type="int", low=1, high=4), | ||
| "e_layers": SearchRange(param_type="int", low=1, high=4), | ||
| "num_layers": SearchRange(param_type="int", low=1, high=4), | ||
| "d_ff": SearchRange(param_type="int", low=64, high=2048, log=True), | ||
| "d_model": SearchRange(param_type="int", low=16, high=512, log=True), | ||
| "patch_length": SearchRange( | ||
| param_type="categorical", choices=[1, 2, 4, 8, 12, 16, 24] | ||
| ), | ||
| "moving_avg": SearchRange( | ||
| param_type="categorical", choices=[3, 5, 7, 11, 15, 21, 25] | ||
| ), | ||
| "persistence_weight": SearchRange(param_type="float", low=0.0, high=1.0), | ||
| "factor": SearchRange(param_type="int", low=1, high=10), | ||
| } | ||
|
|
||
|
|
||
| TRAINER_PARAMS = { | ||
| "gradient_clip_val": SearchRange( | ||
| param_type="float", low=0.01, high=100.0, log=True | ||
| ), | ||
| } | ||
|
|
||
| GLOBAL_SEARCH_SPACE = {**UNIVERSAL_PARAMS, **MODEL_PARAMS, **TRAINER_PARAMS} |
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,129 @@ | ||
| """ | ||
| HyperparameterTuner: Centralized, model-agnostic optimizer. | ||
|
|
||
| """ | ||
|
|
||
| import copy | ||
| import os | ||
|
|
||
| from pytorch_forecasting.tuning.search_range import SearchRange | ||
|
|
||
|
|
||
| class HyperparameterTuner: | ||
| """Model-agnostic hyperparameter optimizer for v2 models. | ||
|
|
||
| Parameters | ||
| ---------- | ||
| pkg_cls : type | ||
| The package class (e.g., TimeXer_pkg_v2, DLinear_pkg_v2). | ||
| data : TimeSeries or LightningDataModule | ||
| Training data. | ||
| base_model_cfg : dict, optional | ||
| Fixed model config that won't be tuned. | ||
| base_trainer_cfg : dict, optional | ||
| Fixed trainer config. | ||
| base_datamodule_cfg : dict, optional | ||
| Fixed datamodule config. | ||
| """ | ||
|
|
||
| def __init__( | ||
| self, | ||
| pkg_cls, | ||
| data, | ||
| base_model_cfg=None, | ||
| base_trainer_cfg=None, | ||
| base_datamodule_cfg=None, | ||
| ): | ||
| self.pkg_cls = pkg_cls | ||
| self.data = data | ||
| self.base_model_cfg = base_model_cfg or {} | ||
| self.base_trainer_cfg = base_trainer_cfg or {} | ||
| self.base_datamodule_cfg = base_datamodule_cfg or {} | ||
|
|
||
| def _build_trial_config(self, trial, search_ranges): | ||
| """Convert Optuna trial suggestions into model_cfg and trainer_cfg. | ||
|
|
||
| This is the 'translation layer' between Optuna and Base_pkg. | ||
| """ | ||
| model_cfg = copy.deepcopy(self.base_model_cfg) | ||
| trainer_cfg = copy.deepcopy(self.base_trainer_cfg) | ||
|
|
||
| for param_name, search_range in search_ranges.items(): | ||
| value = search_range.suggest(trial, param_name) | ||
|
|
||
| if param_name == "gradient_clip_val": | ||
| trainer_cfg["gradient_clip_val"] = value | ||
| elif "." in param_name: | ||
| parts = param_name.split(".") | ||
| d = model_cfg | ||
| for part in parts[:-1]: | ||
| d = d.setdefault(part, {}) | ||
| d[parts[-1]] = value | ||
| else: | ||
| model_cfg[param_name] = value | ||
|
|
||
| return model_cfg, trainer_cfg | ||
|
|
||
| def optimize( | ||
| self, | ||
| n_trials=100, | ||
| timeout=3600 * 8, | ||
| max_epochs=20, | ||
| custom_ranges=None, | ||
| study=None, | ||
| direction="minimize", | ||
| ): | ||
| """Run hyperparameter optimization. | ||
|
|
||
| Parameters | ||
| ---------- | ||
| n_trials : int | ||
| Number of Optuna trials. | ||
| timeout : float | ||
| Maximum time in seconds. | ||
| max_epochs : int | ||
| Max training epochs per trial. | ||
| custom_ranges : dict[str, SearchRange], optional | ||
| Override or add to auto-discovered ranges. | ||
| study : optuna.Study, optional | ||
| Existing study to resume. | ||
|
|
||
| Returns | ||
| ------- | ||
| optuna.Study | ||
| The completed study with results. | ||
| """ | ||
|
|
||
| try: | ||
| import optuna | ||
| except ImportError: | ||
| raise ImportError( | ||
| "Optuna is required for hyperparameter tuning. " | ||
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. why aren't we using |
||
| "Please install it with `pip install optuna`" | ||
| ) | ||
|
|
||
| model_cls = self.pkg_cls.get_cls() | ||
| search_ranges = model_cls.get_tuneable_hyperparameters() | ||
|
|
||
| if custom_ranges: | ||
| search_ranges.update(custom_ranges) | ||
|
|
||
| def objective(trial): | ||
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. should this be a private method?
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. I am talking about the method not the file?
Member
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. I just got a little confused, thanks. I'll make
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Please dont use AI blindly, this creates a lot of confusion and wastes time for all of us |
||
| model_cfg, trainer_cfg = self._build_trial_config(trial, search_ranges) | ||
| trainer_cfg.setdefault("max_epochs", max_epochs) | ||
| trainer_cfg.setdefault("enable_progress_bar", False) | ||
|
|
||
| pkg = self.pkg_cls( | ||
| model_cfg=model_cfg, | ||
| trainer_cfg=trainer_cfg, | ||
| datamodule_cfg=self.base_datamodule_cfg, | ||
| ) | ||
| pkg.fit(self.data, save_ckpt=False) | ||
|
|
||
| return pkg.trainer.callback_metrics["val_loss"].item() | ||
|
|
||
| if study is None: | ||
| study = optuna.create_study(direction=direction) | ||
| study.optimize(objective, n_trials=n_trials, timeout=timeout) | ||
|
|
||
| return study | ||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,51 @@ | ||
| """ | ||
| SearchRange: A typed container for hyperparameter search spaces. | ||
|
|
||
| """ | ||
|
|
||
| from dataclasses import dataclass | ||
| from typing import Any | ||
|
|
||
|
|
||
| @dataclass | ||
| class SearchRange: | ||
| """Defines a search range for a single hyperparameter. | ||
|
|
||
| Parameters | ||
| ---------- | ||
| param_type : str | ||
| One of "int", "float", or "categorical". | ||
| low : float or int, optional | ||
| Lower bound (for int/float types). | ||
| high : float or int, optional | ||
| Upper bound (for int/float types). | ||
| choices : list, optional | ||
| Valid choices (for categorical type). | ||
| log : bool, default=False | ||
| If True, sample in log-uniform space. | ||
| Use for params where relative change matters more than absolute | ||
| (e.g., learning_rate: 1e-5 vs 1e-4 is a 10x change). | ||
| step : int, optional | ||
| Step size for integer parameters. | ||
| """ | ||
|
|
||
| param_type: str | ||
| low: float | int | None = None | ||
| high: float | int | None = None | ||
| choices: list[Any] | None = None | ||
| log: bool = False | ||
| step: int | None = None | ||
|
|
||
| def suggest(self, trial, name: str): | ||
| """Ask Optuna to suggest a value for this parameter. | ||
|
|
||
| This bridges our SearchRange to Optuna's trial API. | ||
| """ | ||
| if self.param_type == "int": | ||
| return trial.suggest_int(name, self.low, self.high, log=self.log) | ||
| elif self.param_type == "float": | ||
| return trial.suggest_float(name, self.low, self.high, log=self.log) | ||
| elif self.param_type == "categorical": | ||
| return trial.suggest_categorical(name, self.choices) | ||
| else: | ||
| raise ValueError(f"Unknown param_type: {self.param_type}") |
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
what if there is no existing study and the user wants to create a new study here?