diff --git a/deeptab/architectures/mambatab.py b/deeptab/architectures/mambatab.py index a70705f..5aa75df 100644 --- a/deeptab/architectures/mambatab.py +++ b/deeptab/architectures/mambatab.py @@ -110,6 +110,7 @@ def forward(self, *data): x = self.norm_f(x) x = self.embedding_activation(x) + x = self.mamba(x) if self.axis == 1: x = x.squeeze(1) else: diff --git a/tests/test_mambatab_backbone.py b/tests/test_mambatab_backbone.py new file mode 100644 index 0000000..2eeeabb --- /dev/null +++ b/tests/test_mambatab_backbone.py @@ -0,0 +1,36 @@ +"""Regression test: MambaTab must actually run its Mamba block. + +forward() went initial_layer -> norm -> activation -> head and never called +self.mamba, so the model trained as a linear layer with an MLP head while all +Mamba parameters sat in the optimizer with no gradients. +""" + +from typing import Any, cast + +import numpy as np +import pandas as pd + +from deeptab.models import MambaTabRegressor + + +def _backward_one_batch(model) -> Any: + task = cast(Any, model._task_model) + task.zero_grad() + (num, cat, emb), labels = next(iter(cast(Any, model._data_module).train_dataloader())) + preds = task(num, cat, emb) + task.compute_loss(preds, labels).backward() + return task + + +def test_mambatab_trains_its_mamba_block(): + rng = np.random.RandomState(0) + X = pd.DataFrame({"num1": rng.randn(60), "num2": rng.rand(60)}) + y = rng.randn(60) + + model = MambaTabRegressor() + model.fit(X, y, max_epochs=1, batch_size=16, accelerator="cpu") + + task = _backward_one_batch(model) + mamba_params = [(name, p) for name, p in task.named_parameters() if ".mamba." in name] + assert mamba_params, "expected mamba parameters on the estimator" + assert all(p.grad is not None for _, p in mamba_params)