Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
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
3 changes: 2 additions & 1 deletion deeptab/training/lightning_module.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"])

Expand Down
19 changes: 15 additions & 4 deletions deeptab/training/optimizers.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
47 changes: 47 additions & 0 deletions tests/test_lightning_module_loss.py
Original file line number Diff line number Diff line change
@@ -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)
41 changes: 40 additions & 1 deletion tests/test_training_optimizers.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand Down
Loading