Skip to content
Closed
Show file tree
Hide file tree
Changes from 3 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
41 changes: 41 additions & 0 deletions tests/models/nemotron_h/test_modeling_nemotron_h.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
# limitations under the License.
"""Testing suite for the PyTorch NemotronH model."""

import copy
import tempfile
import unittest

Expand Down Expand Up @@ -356,6 +357,41 @@ def create_and_check_nemotron_h_chunked_prefill(self, config, input_ids, *args,
msg=f"Max diff: {(ref_first - under_test_first).abs().max().item():.6f}",
)

def create_and_check_nemotron_h_slow_path_multi_chunk(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.

But I'm still a bit confused - why do we have this new test and not just the slow vs fast path test? (which should catch the same regression)

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.

slow_vs_fast can't catch this one: it runs seq_length=7 < chunk_size, i.e. a single chunk with an empty cache, so the inter-chunk recurrence is never reached (previous_states=0, nothing to mix). Buggy and fixed give identical outputs there. It passes unchanged on the buggy code. The regression only shows with more than 2 chunks, hence the separate test.

"""
Regression test for the inter-chunk recurrence of the Mamba2 slow (torch) path: a
single chunked forward over a multi-chunk sequence must reproduce a token-by-token
recurrent decode.

The mixer is fed a large-magnitude input so the SSM state is O(1). With a normal
activation scale the inter-chunk contribution is ~1e-7, i.e. below any tolerance,
which is why neither `create_and_check_mamba2_slow_vs_fast_forward` (single chunk,
needs the kernels) nor a plain `prepare_config_and_inputs` input surfaces the bug.
A small `chunk_size` lets a short sequence span several chunks. Runs on CPU without
the fast-path kernels.
"""
config = copy.deepcopy(config)
config.chunk_size = 8
model = NemotronHModel(config).eval().to(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 mean we need to guarantee that this uses the slow path and here it could result into the fast path. Wouldnt it make more sense to just force cpu here?

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.

Done. The model is now .to("cpu") so it always takes the torch_forward slow path regardless of kernel availability.

mixer = next(layer.mixer for layer in model.layers if getattr(layer, "block_type", None) == "linear_attention")

seq_len = 4 * config.chunk_size + 1
torch.manual_seed(0)
hidden_states = 100.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 mean we are still manually creating here, no? The idea was to directly use the created input ids

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.

OK, thanks for clarifying. It should be fixed now.
feeds model.embeddings(input_ids) (the prepare_config_and_inputs ids) instead of a manual tensor. It's only rescaled so the SSM state is O(1), otherwise the inter-chunk term is ~1e-7 and invisible.


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.parent.assertLess(
max_diff, 1e-3, f"slow-path chunked forward disagrees with recurrent decode: {max_diff}"
)

def prepare_config_and_inputs_for_common(self):
config_and_inputs = self.prepare_config_and_inputs()
(
Expand Down Expand Up @@ -520,6 +556,11 @@ 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.

"""The Mamba2 slow path must reproduce the token-by-token recurrence across chunks."""
config_and_inputs = self.model_tester.prepare_config_and_inputs()
self.model_tester.create_and_check_nemotron_h_slow_path_multi_chunk(*config_and_inputs)

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
59 changes: 59 additions & 0 deletions tests/models/zamba2/test_modeling_zamba2.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
# limitations under the License.
"""Testing suite for the PyTorch Zamba model."""

import copy
import tempfile
import unittest

Expand Down Expand Up @@ -298,6 +299,53 @@ 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.

"""Slow vs fast path check guarded by require kernels to enable fast path"""
model = Zamba2Model(config)
model.eval()
model.to(torch_device)

mamba_mixer = next(layer.mamba for layer in model.layers if hasattr(layer, "mamba"))
hidden_states = model.embed_tokens(input_ids)
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 create_and_check_zamba2_slow_path_multi_chunk(self, config, input_ids, *args):
"""
Regression test for the inter-chunk recurrence of the Mamba2 slow (torch) path: a
single chunked forward over a multi-chunk sequence must reproduce a token-by-token
recurrent decode.

The mixer is fed a large-magnitude input so the SSM state is O(1). With a normal
activation scale the inter-chunk contribution is ~1e-7, i.e. below any tolerance,
which is why neither `create_and_check_zamba2_slow_vs_fast_forward` (single chunk,
needs the kernels) nor a plain `prepare_config_and_inputs` input surfaces the bug.
A small `chunk_size` lets a short sequence span several chunks. Runs on CPU without
the fast-path kernels.
"""
config = copy.deepcopy(config)
config.chunk_size = 8
model = Zamba2Model(config).eval().to(torch_device)
mixer = next(layer.mamba for layer in model.layers if hasattr(layer, "mamba"))

seq_len = 4 * config.chunk_size + 1
torch.manual_seed(0)
hidden_states = 100.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.parent.assertLess(
max_diff, 1e-3, f"slow-path chunked forward disagrees with recurrent decode: {max_diff}"
)

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

@require_torch_accelerator
@require_kernels
def test_mamba2_slow_vs_fast_forward(self):
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):
"""The Mamba2 slow path must reproduce the token-by-token recurrence across chunks."""
config_and_inputs = self.model_tester.prepare_config_and_inputs()
self.model_tester.create_and_check_zamba2_slow_path_multi_chunk(*config_and_inputs)

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