diff --git a/deeptab/training/lightning_module.py b/deeptab/training/lightning_module.py index 9619add..4be87b0 100644 --- a/deeptab/training/lightning_module.py +++ b/deeptab/training/lightning_module.py @@ -240,7 +240,8 @@ def __init__( if not self.loss_fct: self.loss_fct = nn.CrossEntropyLoss() else: - self.loss_fct = nn.MSELoss() + if not self.loss_fct: + self.loss_fct = nn.MSELoss() self.save_hyperparameters(ignore=["model_class", "loss_fct", "family"]) diff --git a/deeptab/training/optimizers.py b/deeptab/training/optimizers.py index d0895e9..ff60a14 100644 --- a/deeptab/training/optimizers.py +++ b/deeptab/training/optimizers.py @@ -297,20 +297,31 @@ def normalize_optimizer_kwargs(optimizer_args: dict[str, Any] | None) -> dict[st >>> normalize_optimizer_kwargs(None) {} - >>> normalize_optimizer_kwargs({"lr": 1e-3}) # non-prefixed key is dropped + >>> normalize_optimizer_kwargs({"eps": 1e-6}) # unprefixed keys pass through + {'eps': 1e-06} + + >>> normalize_optimizer_kwargs({"lr": 1e-3}) # lr/weight_decay are reserved {} Notes ----- + ``lr`` and ``weight_decay`` (prefixed or not) are dropped because + :func:`build_optimizer` receives them as explicit arguments; forwarding + them here would raise a duplicate-keyword error. + This function is called automatically by ``TaskModel.__init__``. You only need to call it directly when building an optimizer outside of ``TaskModel``, e.g. in a custom training loop. """ if not optimizer_args: return {} - return { - key.removeprefix("optimizer_"): value for key, value in optimizer_args.items() if key.startswith("optimizer_") - } + normalized: dict[str, Any] = {} + for key, value in optimizer_args.items(): + key = key.removeprefix("optimizer_") + if key in ("lr", "weight_decay"): + continue + normalized[key] = value + return normalized def build_parameter_groups( diff --git a/tests/test_lightning_module_loss.py b/tests/test_lightning_module_loss.py new file mode 100644 index 0000000..6e5142c --- /dev/null +++ b/tests/test_lightning_module_loss.py @@ -0,0 +1,47 @@ +"""Regression tests for loss selection in TaskModel.""" + +import torch.nn as nn + +from deeptab.configs import MLPConfig +from deeptab.training.lightning_module import TaskModel + + +class _DummyEstimator(nn.Module): + def __init__(self, config=None, feature_information=None, num_classes=1, lss=False, **kwargs): + super().__init__() + self.linear = nn.Linear(4, num_classes) + + def forward(self, *data): + return self.linear(data[0]) + + +def _make_task_model(**overrides): + kwargs = { + "model_class": _DummyEstimator, + "config": MLPConfig(), + "feature_information": ({}, {}, {}), + "num_classes": 1, + } + kwargs.update(overrides) + return TaskModel(**kwargs) + + +class TestLossSelection: + def test_regression_custom_loss_respected(self): + """Regression test: num_classes=1 must not overwrite a custom loss with MSE.""" + custom = nn.HuberLoss() + task = _make_task_model(loss_fct=custom) + assert task.loss_fct is custom + + def test_regression_defaults_to_mse(self): + task = _make_task_model() + assert isinstance(task.loss_fct, nn.MSELoss) + + def test_binary_custom_loss_respected(self): + custom = nn.BCEWithLogitsLoss() + task = _make_task_model(num_classes=2, loss_fct=custom) + assert task.loss_fct is custom + + def test_multiclass_defaults_to_cross_entropy(self): + task = _make_task_model(num_classes=3) + assert isinstance(task.loss_fct, nn.CrossEntropyLoss) diff --git a/tests/test_training_optimizers.py b/tests/test_training_optimizers.py index 6271f87..d9222d0 100644 --- a/tests/test_training_optimizers.py +++ b/tests/test_training_optimizers.py @@ -166,7 +166,7 @@ def test_strips_prefix(self): assert result == {"betas": (0.9, 0.95)} def test_non_prefixed_keys_excluded(self): - # Only keys that START with "optimizer_" are kept + # lr/weight_decay are reserved (passed explicitly by build_optimizer) result = normalize_optimizer_kwargs({"optimizer_eps": 1e-8, "lr": 1e-3}) assert "eps" in result assert "lr" not in result @@ -176,6 +176,45 @@ def test_multiple_keys(self): result = normalize_optimizer_kwargs(raw) assert result == {"betas": (0.9, 0.99), "eps": 1e-8} + def test_unprefixed_keys_pass_through(self): + """Regression test: TrainerConfig.optimizer_kwargs uses unprefixed keys. + + They were previously filtered out entirely, so user optimizer settings + never reached the optimizer constructor. + """ + result = normalize_optimizer_kwargs({"eps": 1e-1, "betas": (0.5, 0.6)}) + assert result == {"eps": 1e-1, "betas": (0.5, 0.6)} + + def test_reserved_keys_dropped_regardless_of_prefix(self): + result = normalize_optimizer_kwargs({"optimizer_lr": 1e-3, "weight_decay": 0.1, "eps": 1e-8}) + assert result == {"eps": 1e-8} + + def test_trainer_config_optimizer_kwargs_reach_the_optimizer(self): + """End-to-end regression test for the TrainerConfig path.""" + from typing import Any, cast + + import numpy as np + import pandas as pd + + from deeptab.configs import TrainerConfig + from deeptab.models import MLPRegressor + + rng = np.random.RandomState(0) + X = pd.DataFrame({"a": rng.randn(40)}) + y = rng.randn(40) + model = MLPRegressor( + trainer_config=TrainerConfig( + max_epochs=1, + optimizer_type="AdamW", + optimizer_kwargs={"eps": 1e-1, "betas": (0.5, 0.6)}, + ) + ) + model.fit(X, y, max_epochs=1, batch_size=16, accelerator="cpu") + task_model = cast(Any, model._task_model) + group = task_model.trainer.optimizers[0].param_groups[0] + assert group["eps"] == 1e-1 + assert group["betas"] == (0.5, 0.6) + # --------------------------------------------------------------------------- # build_parameter_groups