From e66d8e29bb23792deb330dccb48129d676533ee8 Mon Sep 17 00:00:00 2001 From: Yuhe Zhang Date: Mon, 21 Sep 2026 07:35:27 -0700 Subject: [PATCH] fix(nemotron_h): honor configured Mamba time-step limits Signed-off-by: Yuhe Zhang --- .../models/nemotron_h/modeling_nemotron_h.py | 3 +- .../models/nemotron_h/modular_nemotron_h.py | 1 + .../nemotron_h/test_modeling_nemotron_h.py | 44 +++++++++++++++++++ 3 files changed, 46 insertions(+), 2 deletions(-) diff --git a/src/transformers/models/nemotron_h/modeling_nemotron_h.py b/src/transformers/models/nemotron_h/modeling_nemotron_h.py index f21ca65ee59a..7484caf00bda 100644 --- a/src/transformers/models/nemotron_h/modeling_nemotron_h.py +++ b/src/transformers/models/nemotron_h/modeling_nemotron_h.py @@ -377,8 +377,7 @@ def __init__(self, config: NemotronHConfig, layer_idx: int | None = None, initia self.n_groups = config.n_groups self.head_dim = config.mamba_head_dim self.chunk_size = config.chunk_size - # No upper limit - self.time_step_limit = (config.time_step_min, float("inf")) + self.time_step_limit = config.time_step_limit self.conv_dim = self.intermediate_size + 2 * self.n_groups * self.ssm_state_size diff --git a/src/transformers/models/nemotron_h/modular_nemotron_h.py b/src/transformers/models/nemotron_h/modular_nemotron_h.py index af9b51d72fdf..fa35a0937c21 100644 --- a/src/transformers/models/nemotron_h/modular_nemotron_h.py +++ b/src/transformers/models/nemotron_h/modular_nemotron_h.py @@ -58,6 +58,7 @@ def __init__(self, config: NemotronHConfig, layer_idx: int | None = None, initia self.n_groups = config.n_groups self.head_dim = config.mamba_head_dim self.num_heads = config.mamba_num_heads + self.time_step_limit = config.time_step_limit self.conv1d = nn.Conv1d( in_channels=self.conv_dim, diff --git a/tests/models/nemotron_h/test_modeling_nemotron_h.py b/tests/models/nemotron_h/test_modeling_nemotron_h.py index 3786fd80fa41..605945361488 100644 --- a/tests/models/nemotron_h/test_modeling_nemotron_h.py +++ b/tests/models/nemotron_h/test_modeling_nemotron_h.py @@ -13,6 +13,7 @@ # limitations under the License. """Testing suite for the PyTorch NemotronH model.""" +import copy import tempfile import unittest @@ -42,6 +43,7 @@ import torch from transformers import DynamicCache, NemotronHForCausalLM, NemotronHModel, StaticCache + from transformers.models.nemotron_h.modeling_nemotron_h import NemotronHMamba2Mixer class NemotronHModelTester: @@ -529,6 +531,48 @@ def test_model(self): config_and_inputs = self.model_tester.prepare_config_and_inputs() self.model_tester.create_and_check_model(*config_and_inputs) + def test_mamba_time_step_min_only_affects_initialization(self): + torch.manual_seed(0) + config = self.model_tester.get_config() + mixer = NemotronHMamba2Mixer(config, layer_idx=0).eval() + with torch.no_grad(): + mixer.dt_bias.fill_(-12.0) + mixer.D.zero_() + + other_config = copy.deepcopy(config) + other_config.time_step_min = 0.1 + other_mixer = NemotronHMamba2Mixer(other_config, layer_idx=0).eval() + other_mixer.load_state_dict(mixer.state_dict()) + + hidden_states = torch.randn(2, 7, config.hidden_size, requires_grad=True) + output = mixer(hidden_states) + other_output = other_mixer(hidden_states) + torch.testing.assert_close(output, other_output) + gradient = torch.autograd.grad(output.square().sum(), hidden_states)[0] + other_gradient = torch.autograd.grad(other_output.square().sum(), hidden_states)[0] + torch.testing.assert_close(gradient, other_gradient) + + def test_mamba_time_step_limit_bounds(self): + torch.manual_seed(0) + config = self.model_tester.get_config() + config.time_step_limit = (0.01, 0.02) + mixer = NemotronHMamba2Mixer(config, layer_idx=0).eval() + with torch.no_grad(): + mixer.in_proj.weight[-config.mamba_num_heads :].zero_() + mixer.D.zero_() + reference_config = copy.deepcopy(config) + reference_config.time_step_limit = (0.0, float("inf")) + reference = NemotronHMamba2Mixer(reference_config, layer_idx=0).eval() + reference.load_state_dict(mixer.state_dict()) + hidden_states = torch.randn(2, 7, config.hidden_size) + + for dt_bias, expected_dt in [(-12.0, 0.01), (12.0, 0.02)]: + with self.subTest(dt_bias=dt_bias), torch.no_grad(): + mixer.dt_bias.fill_(dt_bias) + dt = torch.tensor(expected_dt) + reference.dt_bias.copy_((dt + torch.log(-torch.expm1(-dt))).expand_as(reference.dt_bias)) + torch.testing.assert_close(mixer(hidden_states), reference(hidden_states)) + def test_for_causal_lm(self): config_and_inputs = self.model_tester.prepare_config_and_inputs() self.model_tester.create_and_check_for_causal_lm(*config_and_inputs)