Skip to content
Open
Show file tree
Hide file tree
Changes from 3 commits
Commits
Show all changes
38 commits
Select commit Hold shift + click to select a range
bb7d41c
feat : Implement a unified hyperparameter tuning interface decoupled …
Muhammad-Rebaal Jul 3, 2026
0c691db
feat : Implement a unified hyperparameter tuning interface decoupled …
Muhammad-Rebaal Jul 7, 2026
772fa91
fix pytest
Muhammad-Rebaal Jul 7, 2026
8301098
feat: Implemented Design wrapper similar to sktime
Muhammad-Rebaal Jul 12, 2026
5133352
Merge branch 'main' into optimization
Muhammad-Rebaal Jul 14, 2026
bac8e87
passed the cfgs to class,added test and adjust the param_grid
Muhammad-Rebaal Jul 15, 2026
8a8f97b
Merge branch 'main' into optimization
Muhammad-Rebaal Jul 21, 2026
4dcd61c
ENH: Tags are taken _BasePtForecasterV2, lightening tuner functions a…
Muhammad-Rebaal Jul 22, 2026
cdd7f85
Merge branch 'optimization' of https://github.com/Muhammad-Rebaal/pyt…
Muhammad-Rebaal Jul 22, 2026
39d5122
ENH: Tags are taken _BasePtForecasterV2, lightening tuner functions a…
Muhammad-Rebaal Jul 22, 2026
928b8a4
Merge branch 'main' of https://github.com/Muhammad-Rebaal/pytorch-for…
Muhammad-Rebaal Aug 12, 2026
b6bf90a
Merge branch 'main' of https://github.com/Muhammad-Rebaal/pytorch-for…
Muhammad-Rebaal Aug 22, 2026
4ff83c2
Merge branch 'main' into optimization
Muhammad-Rebaal Aug 22, 2026
f126396
code cleaned
Muhammad-Rebaal Aug 22, 2026
5788e74
Merge branch 'optimization' of https://github.com/Muhammad-Rebaal/pyt…
Muhammad-Rebaal Aug 22, 2026
761d72e
Merge branch 'main' into optimization
Muhammad-Rebaal Aug 22, 2026
bed2247
HyperparameterTuner Feature Implemented
Muhammad-Rebaal Aug 26, 2026
3c8d8d1
Fix Code Quality
Muhammad-Rebaal Aug 26, 2026
786ac9a
Fix Soft Depency issue
Muhammad-Rebaal Aug 26, 2026
52b28ab
Merge branch 'main' into optimization
Muhammad-Rebaal Aug 26, 2026
1470d12
HyperparameterTuner implemented without the need of datamodule_cfg an…
Muhammad-Rebaal Aug 30, 2026
c341a12
Merge branch 'main' into optimization
Muhammad-Rebaal Aug 31, 2026
ee7f848
Added an edge case
Muhammad-Rebaal Aug 31, 2026
d02820a
Merge branch 'optimization' of https://github.com/Muhammad-Rebaal/pyt…
Muhammad-Rebaal Aug 31, 2026
7d03d39
datamodule_cls variable added into the HyperparameterTuner Class
Muhammad-Rebaal Sep 8, 2026
c8f13d5
Merge branch 'main' into optimization
Muhammad-Rebaal Sep 8, 2026
a2a2a7e
Merge branch 'main' into optimization
Muhammad-Rebaal Sep 10, 2026
2eaa29c
Merge branch 'main' into optimization
Muhammad-Rebaal Sep 11, 2026
87f018e
Merge branch 'main' into optimization
Muhammad-Rebaal Sep 14, 2026
f7662c3
Merge branch 'main' into optimization
Muhammad-Rebaal Sep 16, 2026
300ec2b
cleaned up the test file and added the example
Muhammad-Rebaal Sep 16, 2026
f86629c
Merge branch 'optimization' of https://github.com/Muhammad-Rebaal/pyt…
Muhammad-Rebaal Sep 16, 2026
51e328c
Merge branch 'main' into optimization
Muhammad-Rebaal Sep 18, 2026
51fe98f
Merge branch 'main' into optimization
Muhammad-Rebaal Sep 25, 2026
78fd85c
Merge branch 'main' into optimization
Muhammad-Rebaal Sep 28, 2026
f838efd
Merge branch 'main' into optimization
Muhammad-Rebaal Sep 29, 2026
076338c
Merge branch 'main' into optimization
Muhammad-Rebaal Oct 10, 2026
7a10b01
Added Hparam tuning for b_model params
Muhammad-Rebaal Oct 11, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
55 changes: 55 additions & 0 deletions pytorch_forecasting/tuning/global_registry.py
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}
129 changes: 129 additions & 0 deletions pytorch_forecasting/tuning/hyperparameter_tuner.py
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.

Copy link
Copy Markdown
Member

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?


Returns
-------
optuna.Study
The completed study with results.
"""

try:
import optuna
except ImportError:
raise ImportError(
"Optuna is required for hyperparameter tuning. "

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

why aren't we using _safe_import or _check_soft_dependencies here?

"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):

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

should this be a private method?

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I am talking about the method not the file?
And this is not even the file that I have added this comment on?

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I just got a little confused, thanks. I'll make objective method private as its an internal one.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The 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
51 changes: 51 additions & 0 deletions pytorch_forecasting/tuning/search_range.py
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}")
Loading