Repository navigation
[ENH] Implement a unified hyperparameter tuning interface decoupled from model classes for v2 - #2335
Muhammad-Rebaal wants to merge 37 commits into
Conversation
…from model classes for v2
…from model classes for v2
Codecov Report❌ Patch coverage is
Additional details and impacted files@@ Coverage Diff @@
## main #2335 +/- ##
=======================================
Coverage ? 88.85%
=======================================
Files ? 219
Lines ? 11772
Branches ? 0
=======================================
Hits ? 10460
Misses ? 1312
Partials ? 0
Flags with carried forward coverage won't be shown. Click here to find out more. ☔ View full report in Codecov by Harness. 🚀 New features to boost your workflow:
|
phoeenniixx
left a comment
There was a problem hiding this comment.
Why aren't we using optuna here? Also, in the param grid the values are hard coded
Because of these 2 reason : User Experience: It's much easier for a user to inspect, modify, and pass a simple SearchRange dataclass in custom_ranges than it is to write and pass custom Optuna lambda functions. Decoupling: By keeping Optuna logic out of the base model layer and registry, the models themselves remain completely framework-agnostic. The SearchRange dataclass acts as a clean bridge.
In v1's optimize_hyperparameters for TFT, the default search spaces were hardcoded as tuple arguments (e.g., hidden_size_range=(16, 265), dropout_range=(0.1, 0.3)). The global_registry simply takes those v1 hardcoded defaults and standardizes them centrally so they can apply to all models automatically. However, they aren't strictly locked users can completely override these default "hardcoded" values by passing their own custom_ranges to the HyperparameterTuner. |
fkiraly
left a comment
There was a problem hiding this comment.
Hm, ok start. I would make all classes and variables private for now, so we can later change. This increases separation of concerns which is good.
Some questions in the design:
- how would this get used?
- This designs custom search range classes. Do we want to do that, instead of relying on one of the standards? (e.g., even
optunaitself) - optuna is one choice of what we could do. Have a look at the
hyperactivetuning class for neural networks, or theOptCVclass insktime. Could this be useful?
|
Thanks for the review.
I've implemented this.
Yes, it is preciously useful, helped me a lot and implemented the wrapper class |
There was a problem hiding this comment.
- Please use soft-dep isolation when importing soft dependencies
See https://www.sktime.net/docs/developer-guide/dependencies/#isolating-soft-dependencies-to-estimators - I think
lightning.tunercan also be very helpful here, no? - Please add tests
| try: | ||
| import optuna | ||
| except ImportError: | ||
| raise ImportError( | ||
| "Optuna is required for hyperparameter tuning. " |
There was a problem hiding this comment.
why aren't we using _safe_import or _check_soft_dependencies here?
| custom_ranges : dict[str, SearchRange], optional | ||
| Override or add to auto-discovered ranges. | ||
| study : optuna.Study, optional | ||
| Existing study to resume. |
There was a problem hiding this comment.
what if there is no existing study and the user wants to create a new study here?
| if custom_ranges: | ||
| search_ranges.update(custom_ranges) | ||
|
|
||
| def objective(trial): |
There was a problem hiding this comment.
should this be a private method?
There was a problem hiding this comment.
I am talking about the method not the file?
And this is not even the file that I have added this comment on?
There was a problem hiding this comment.
I just got a little confused, thanks. I'll make objective method private as its an internal one.
There was a problem hiding this comment.
Please dont use AI blindly, this creates a lot of confusion and wastes time for all of us
based on this vignette: Also, why do we need to pass the cfgs again to this class, I thought we could simply pass the |
Nice Idea I'll try to implement that. |
Yes,
These are the methods I found that can be helpful. Any suggestions from your side.. @phoeenniixx |
…dded code refactored and removed unwanted tests
|
Hi @phoeenniixx, Now the Hyperparameters can be tuned for any v2 model like this : Although the user can provide their |
|
Hi @phoeenniixx,
The user can now pass raw timeseries and datamodule as well if passed raw timeseries datamodule would be made to processed and if lightening datamodule is passed then the user can also specify the train-test-split that we are working on another PR. Lastly, every model has a get_datamodule_cls method implemented in their respective pkg coming from base_pkg, and since we are moving towards deprecation of the base_pkg. I've moved the get_datamodule_cls in their respective model_cls such as
The purpose of this is that the user don't have to specify which datamodule is for which model and converts their data into the specific datamodule to be processed also added the edge case if the user passes the wrong datamodule it would raise a |
|
Hi @phoeenniixx,
|
| or against which to validate when ``data`` is a prebuilt DataModule. | ||
| **fixed_hparams | ||
| Any model parameter that should stay constant, e.g. | ||
| ``hidden_size=128``. |
There was a problem hiding this comment.
please add some examples here to how to use this
phoeenniixx
left a comment
There was a problem hiding this comment.
Thanks
Can you please clean up the test file? it has a lot of "param checks" like if i pass a, it is collected in x.param etc. It checks if something is implemented or not - we should only test if the things are working correctly... Not that if the tests accpet some param or not, it is understood that the class will accept datamodule and assign it to tuner.datamoudle.
(There are multiple tests like that here)
Also, pls add some examples (maybe even an example nb!)
|
Hi @phoeenniixx, |
phoeenniixx
left a comment
There was a problem hiding this comment.
Thanks!
Will this work for a model checkpoint etc as well?
Also, should we also add a nb etc - as rn we cant verify that the complete workflow works as expected
| if datamodule_cls is not None: | ||
| if not ( | ||
| isinstance(datamodule_cls, type) | ||
| and issubclass(datamodule_cls, LightningDataModule) |
There was a problem hiding this comment.
do we need this validation? if we add type hints etc, it would be understood that it is a LightningDataModule no?
| trainer.fit(model, datamodule=self.datamodule) | ||
|
|
||
| metrics = trainer.callback_metrics | ||
| for key in ("val_loss", "train_loss_epoch", "train_loss"): |
There was a problem hiding this comment.
should we optimize for train losses? that might lead to overfitting?
There was a problem hiding this comment.
yes, you're right over here
| custom_ranges = custom_ranges or {} | ||
|
|
||
| for param_name, param_obj in sig.parameters.items(): | ||
| if param_name in base_params or param_name == "self": |
There was a problem hiding this comment.
Here what will happen to the params which are sent to theBaseModel? Like Samformer sends optimizer etc directly to base model:
There was a problem hiding this comment.
Right now it gonna skip, in my mind we're tuning the model specific ones. I think skipping self and logging_metrics instead of whole base_model param would be a cleaner fix
There was a problem hiding this comment.
but should we also skip optimizer, lr_scheduler etc?
There was a problem hiding this comment.
For optimizer_params , lr_scheduler , these 2 have default none it would just auto skip as we've a logic:
if (
default is inspect.Parameter.empty
or default is None
or isinstance(default, (dict, list))
):
continue
-
For
optimizerandloss, I think if we skip the users would lost the ability to pass their custom ranges ? -
Optuna can't tune dictionaries like
optimizer_paramsdirectly anyway. Since the learning rate is universally critical I think we just inject it directly inside the_objective function(just like the old TFT tuning.py did):
if "optimizer_params" not in self.fixed_hparams:
lr = trial.suggest_float("lr", 1e-5, 1e-1, log=True)
model_cfg["optimizer_params"] = {"lr": lr}
There was a problem hiding this comment.
I think i dont understand the issue - i think we need to optimise some of the params atleast - like lr_scheduler, no? So, how should we do that?
Can you please elaborate on this?
There was a problem hiding this comment.
I thought above you're telling me not to tune lr_scheduler that's why I said we already have a logic to skip and suggested what we can do with other params.
There was a problem hiding this comment.
My ques is that how are these params optimized here? Are we optimizing lr_scheduler? And come to think of it, for lr, we can also use tuner from lightning, no? That would be better?
There was a problem hiding this comment.
Sorry, for not being clear :)
I just want to know how are the params that we are passing to base_model directly being optimized?
Fixes : #2332
This PR implements the hyperparameter optimization interface for v2 models.
Implemented Architecture
SearchRange& Global Registry (tuning/search_range.py,tuning/global_registry.py):hidden_size,dropout,n_heads).Auto-Discovery Hook (
models/base/_base_model_v2.py):get_tuneable_hyperparameters()class method onBaseModelthat usesinspect.signature(cls.__init__)to detect what parameters the subclass accepts and maps them to the global registry.HyperparameterTuner(tuning/hyperparameter_tuner.py):Base_pkg.fit().User Flow: