Fix Mamba2 chunked-prefill / speculative decoding for Zamba2, Nemotron-H, Bamba, FalconH1 and GraniteMoeHybrid (#46741)
* Fix Mamba2 mixer chunked-prefill/speculative decode for Zamba2 and Nemotron-H
The Mamba2 mixer took the single-step decode branch whenever the cache
held a previous state, ignoring seq_len. For seq_len > 1 with a non-empty
cache (chunked prefill, speculative/assisted decode verification) it used
only the first tokens dt plus one conv/SSM step, silently producing wrong
logits for the rest of the chunk and corrupting the cached state.
Gate the single-step branch on seq_len == 1; for the multi-token cached
case run the full path seeded from the cache (prepend the cached conv
left context, pass the cached recurrent state as the initial SSM state).
Same root cause and fix as mamba2 (#46032); applied in modular_zamba2.py
so the change propagates to nemotron_h, which inherits the mixer.
Signed-off-by: Ting Sun <suntcrick@gmail.com>
* Prepend the full cached conv_states in the Mamba2 chunked path
Match mamba2 (#46084) and the qwen3-next reference by prepending the full
cached conv_states instead of conv_states[:, :, 1:]. The two are numerically
identical: the extra oldest column only adds one warm-up conv position that the
trailing [-seq_len:] slice drops. This keeps the chunked-prefill left-context
consistent with the reference path.
Signed-off-by: Ting Sun <suntcrick@gmail.com>
* Fix Mamba2 cuda_kernels_forward chunked-prefill for Zamba2 and Nemotron-H
cuda_kernels_forward had the same defect as the torch path: a multi-token
forward with a primed cache (chunked prefill / speculative verification) was
routed into the single-step causal_conv1d_update / selective_state_update
kernels and crashed ("weight must have shape (dim, width)"). Gate the
single-step branch on seq_len == 1; for the multi-token cached case run the
full kernels seeded from the cache (prepend the cached conv left-context and
drop it after the conv, pass the cached recurrent state as initial_states to
mamba_chunk_scan_combined), mirroring #46084.
Signed-off-by: Ting Sun <suntcrick@gmail.com>
* Align Zamba2/Nemotron-H chunked-prefill tests with merged mamba2 test
Reshape the regression test to match the form merged in #46084: a model-level
first-token causal check (the first position of a multi-token cached forward must
equal the single-token cached forward), with separate CPU and accelerator variants.
Nemotron-H has no CPU variant because its MoE layers use grouped_mm, which has no
CPU kernel.
Signed-off-by: Ting Sun <suntcrick@gmail.com>
* Propagate Mamba2 chunked-prefill cache fix to Bamba, FalconH1, GraniteMoeHybrid
These models carry their own Mamba2 mixer and gate the single-step branch on
seq_len == 1, but their multi-token cached path still starts from a zero conv
context and zero SSM state, so chunked-prefill / speculative verification is
history-less. Seed the multi-token cached path from the cache (prepend the cached
conv left-context, pass the cached recurrent state as the SSM initial state),
mirroring the Zamba2/Nemotron-H fix and #46084. Add the chunked-prefill regression
test to each.
Signed-off-by: Ting Sun <suntcrick@gmail.com>
* Extract cached conv/recurrent state into locals in the Mamba2 mixers
Bind the cached conv_states / recurrent_states once after use_precomputed_states, matching mamba2, instead of re-reading cache_params.layers[self.layer_idx] inline in the cat / initial_states / single-step paths. Pure refactor, no behavior change.
Signed-off-by: Ting SUN <suntcrick@gmail.com>
* Add CPU variant for the nemotron_h chunked-prefill test
NemotronH MoE experts dispatch through integrations.moe to torch._grouped_mm, which only gained a CPU kernel in torch 2.11; guard the new _cpu test with require_torch_greater_or_equal("2.11") so it runs on current CI and skips on older torch.
Signed-off-by: Ting SUN <suntcrick@gmail.com>
* rename a bit
---------
Signed-off-by: Ting Sun <suntcrick@gmail.com>
Signed-off-by: Ting SUN <suntcrick@gmail.com>
Co-authored-by: vasqu <antonprogamer@gmail.com>