transformers
fdb6d318 - Fix Mamba2 chunked-prefill / speculative decoding for Zamba2, Nemotron-H, Bamba, FalconH1 and GraniteMoeHybrid (#46741)

Commit
91 days ago
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>
Author
Parents
Loading