Skip to content
Closed
Show file tree
Hide file tree
Changes from 2 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
5 changes: 2 additions & 3 deletions src/transformers/models/nemotron_h/modeling_nemotron_h.py
Original file line number Diff line number Diff line change
Expand Up @@ -551,9 +551,8 @@ def torch_forward(self, input_states, cache_params: Cache | None=None, attention
states = torch.cat([previous_states, states], dim=1)
decay_chunk = torch.exp(segment_sum(nn.functional.pad(A_cumsum[:, :, :, -1], (1, 0))))

states_permuted = states.permute(0, 2, 1, 3, 4)
result = (decay_chunk[..., None, None] * states_permuted[:, :, None, ...]).sum(dim=2)
new_states = result.permute(0, 2, 1, 3, 4)
decay_chunk = decay_chunk.transpose(1, 3)
new_states = (decay_chunk[..., None, None] * states[:, :, None, ...]).sum(dim=1)
states, ssm_state = new_states[:, :-1], new_states[:, -1]

# Compute state -> output conversion per chunk
Expand Down
5 changes: 2 additions & 3 deletions src/transformers/models/zamba2/modeling_zamba2.py
Original file line number Diff line number Diff line change
Expand Up @@ -839,9 +839,8 @@ def torch_forward(self, input_states, cache_params: Cache | None=None, attention
states = torch.cat([previous_states, states], dim=1)
decay_chunk = torch.exp(segment_sum(nn.functional.pad(A_cumsum[:, :, :, -1], (1, 0))))

states_permuted = states.permute(0, 2, 1, 3, 4)
result = (decay_chunk[..., None, None] * states_permuted[:, :, None, ...]).sum(dim=2)
new_states = result.permute(0, 2, 1, 3, 4)
decay_chunk = decay_chunk.transpose(1, 3)
new_states = (decay_chunk[..., None, None] * states[:, :, None, ...]).sum(dim=1)
states, ssm_state = new_states[:, :-1], new_states[:, -1]

# Compute state -> output conversion per chunk
Expand Down
5 changes: 2 additions & 3 deletions src/transformers/models/zamba2/modular_zamba2.py
Original file line number Diff line number Diff line change
Expand Up @@ -627,9 +627,8 @@ def torch_forward(self, input_states, cache_params: Cache | None=None, attention
states = torch.cat([previous_states, states], dim=1)
decay_chunk = torch.exp(segment_sum(nn.functional.pad(A_cumsum[:, :, :, -1], (1, 0))))

states_permuted = states.permute(0, 2, 1, 3, 4)
result = (decay_chunk[..., None, None] * states_permuted[:, :, None, ...]).sum(dim=2)
new_states = result.permute(0, 2, 1, 3, 4)
decay_chunk = decay_chunk.transpose(1, 3)
new_states = (decay_chunk[..., None, None] * states[:, :, None, ...]).sum(dim=1)
states, ssm_state = new_states[:, :-1], new_states[:, -1]

# Compute state -> output conversion per chunk
Expand Down
42 changes: 42 additions & 0 deletions tests/models/nemotron_h/test_modeling_nemotron_h.py
Original file line number Diff line number Diff line change
Expand Up @@ -520,6 +520,48 @@ def test_mamba2_slow_vs_fast_forward(self):
config_and_inputs = self.model_tester.prepare_config_and_inputs()
self.model_tester.create_and_check_mamba2_slow_vs_fast_forward(*config_and_inputs)

def test_mamba2_slow_path_multi_chunk(self):

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Hmm was create_and_check_mamba2_slow_vs_fast_forward not enough? Or why the whole new test design?

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks. I reworked into a create_and_check_* helper driven by prepare_config_and_inputs.

"""
Regression test for the inter-chunk recurrence in the Mamba2 slow (torch) path.

A single chunked forward over a multi-chunk sequence must match a token-by-token
recurrent decode. The input is deliberately large-magnitude so the SSM state is
O(1) and the inter-chunk contribution is not numerically negligible; with the
previous `.sum(dim=2)` reduction the two disagree by orders of magnitude. Runs on
CPU without the fast-path kernels.
"""
config = NemotronHConfig(
vocab_size=99,
hidden_size=32,
mamba_num_heads=8,
mamba_head_dim=8,
ssm_state_size=16,
n_groups=1,
mamba_chunk_size=8,
num_attention_heads=2,
num_key_value_heads=2,
head_dim=8,
intermediate_size=32,
use_mamba_kernels=False,
layers_block_type=["mamba"],
)
torch.manual_seed(0)
mixer = NemotronHModel(config).eval().to(torch_device).layers[0].mixer

seq_len = 5 * config.chunk_size + 3
hidden_states = 50.0 * torch.randn(1, seq_len, config.hidden_size, device=torch_device)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I would like to use prepare config and inputs if possible (if this test is still really needed after my comment above)


with torch.no_grad():
chunked = mixer.torch_forward(hidden_states)
cache = DynamicCache(config=config)
recurrent = torch.cat(
[mixer.torch_forward(hidden_states[:, t : t + 1], cache_params=cache) for t in range(seq_len)],
dim=1,
)

max_diff = (chunked - recurrent).abs().max().item()
self.assertLess(max_diff, 1e-3, f"slow-path chunked forward disagrees with recurrent decode: {max_diff}")

def test_attention_outputs(self):
r"""
Overriding the test_attention_outputs test as the NemotronH model outputs attention only for its attention layers
Expand Down
78 changes: 78 additions & 0 deletions tests/models/zamba2/test_modeling_zamba2.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,7 @@
slow,
torch_device,
)
from transformers.utils.import_utils import is_causal_conv1d_available, is_mamba_ssm_available

from ...generation.test_utils import GenerationTesterMixin
from ...test_configuration_common import ConfigTester
Expand Down Expand Up @@ -298,6 +299,34 @@ def create_and_check_zamba2_chunked_prefill(self, config, input_ids, *args, devi
msg=f"Max diff: {(ref_first - under_test_first).abs().max().item():.6f}",
)

def create_and_check_zamba2_slow_vs_fast_forward(self, config, input_ids, *args):

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think it changes a bit on main re mamba2 especially re decorators can you check

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Zamba2 slow_vs_fast now gated with @require_torch_accelerator + @require_kernels, helper trimmed to match mamba2.

"""
Test that cuda_kernels_forward and torch_forward produce consistent outputs for the
Mamba2 mixer, i.e. that the optimized CUDA kernel path and the pure PyTorch path are
equivalent. Guarded by the availability of the fast-path kernels and a CUDA device.
"""
if not (is_mamba_ssm_available() and is_causal_conv1d_available()):
self.parent.skipTest(
"This test needs the Mamba2 fast path. Skipping as the necessary packages have not been found."
)
if torch_device != "cuda":
self.parent.skipTest("This test needs the Mamba2 fast path. Skipping as we need a cuda capable device.")

model = Zamba2Model(config)
model.eval()
model.to(torch_device)

# Find the first Mamba mixer in the model
mamba_mixer = next((layer.mamba for layer in model.layers if hasattr(layer, "mamba")), None)
if mamba_mixer is None:
self.parent.skipTest("No mamba layer found in the model configuration.")

hidden_states = model.embed_tokens(input_ids.to(torch_device))

outputs_fast = mamba_mixer.cuda_kernels_forward(hidden_states)
outputs_slow = mamba_mixer.torch_forward(hidden_states)
self.parent.assertTrue(torch.allclose(outputs_fast, outputs_slow, atol=1e-3, rtol=1e-3))

def prepare_config_and_inputs_for_common(self):
config_and_inputs = self.prepare_config_and_inputs()
(
Expand Down Expand Up @@ -351,6 +380,55 @@ def setUp(self):
self.model_tester = Zamba2ModelTester(self)
self.config_tester = ConfigTester(self, config_class=Zamba2Config, hidden_size=32)

def test_mamba2_slow_vs_fast_forward(self):
"""
Test that cuda_kernels_forward and torch_forward produce consistent outputs.
"""
config_and_inputs = self.model_tester.prepare_config_and_inputs()
self.model_tester.create_and_check_zamba2_slow_vs_fast_forward(*config_and_inputs)

def test_mamba2_slow_path_multi_chunk(self):
"""
Regression test for the inter-chunk recurrence in the Mamba2 slow (torch) path.

A single chunked forward over a multi-chunk sequence must match a token-by-token
recurrent decode. The input is deliberately large-magnitude so the SSM state is
O(1) and the inter-chunk contribution is not numerically negligible; with the
previous `.sum(dim=2)` reduction the two disagree by orders of magnitude. Runs on
CPU without the fast-path kernels.
"""
config = Zamba2Config(
vocab_size=99,
hidden_size=32,
mamba_d_state=16,
num_hidden_layers=1,
num_attention_heads=2,
n_mamba_heads=8,
intermediate_size=8,
chunk_size=8,
mamba_ngroups=1,
use_mamba_kernels=False,
layers_block_type=["mamba"],
num_mem_blocks=1,
use_mem_rope=True,
)
torch.manual_seed(0)
mixer = Zamba2Model(config).eval().to(torch_device).layers[0].mamba

seq_len = 5 * config.chunk_size + 3
hidden_states = 50.0 * torch.randn(1, seq_len, config.hidden_size, device=torch_device)

with torch.no_grad():
chunked = mixer.torch_forward(hidden_states)
cache = DynamicCache(config=config)
recurrent = torch.cat(
[mixer.torch_forward(hidden_states[:, t : t + 1], cache_params=cache) for t in range(seq_len)],
dim=1,
)

max_diff = (chunked - recurrent).abs().max().item()
self.assertLess(max_diff, 1e-3, f"slow-path chunked forward disagrees with recurrent decode: {max_diff}")

@unittest.skip("We need at leat 3 layers to test weight tying!")
def test_num_layers_is_small(self):
pass
Expand Down