[Nemotron] Fix decode track-save reading the stale tail of the CUDA-graph track buffer (#32555)
Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Fable 5
parent
20d5b91e5b
commit
62dfaaa0e0
@@ -566,12 +566,15 @@ class MambaAttnBackendBase(AttentionBackend):
|
|||||||
self.state_indices_list[bs - 1][: len(mamba_indices)].copy_(mamba_indices)
|
self.state_indices_list[bs - 1][: len(mamba_indices)].copy_(mamba_indices)
|
||||||
# Refresh the static track-dest buffer in-place (translated); the captured
|
# Refresh the static track-dest buffer in-place (translated); the captured
|
||||||
# track-save reads it, leaving the handed-in InputBuffer slot read-only.
|
# track-save reads it, leaving the handed-in InputBuffer slot read-only.
|
||||||
|
# Hand out only the refreshed [:bs] prefix — Mamba2's track-save slices
|
||||||
|
# [-num_decodes:], which on the full max_bs buffer binds the stale tail.
|
||||||
track_buf = None
|
track_buf = None
|
||||||
if mamba_track_indices is not None:
|
if mamba_track_indices is not None:
|
||||||
track_buf = self.mamba_track_indices_buf
|
assert (
|
||||||
track_buf[: len(mamba_track_indices)].copy_(
|
len(mamba_track_indices) >= bs
|
||||||
self._translate_mamba_indices(mamba_track_indices)
|
), f"{len(mamba_track_indices)=} < {bs=}"
|
||||||
)
|
track_buf = self.mamba_track_indices_buf[:bs]
|
||||||
|
track_buf.copy_(self._translate_mamba_indices(mamba_track_indices[:bs]))
|
||||||
# Refresh the static write cursor in-place (mirrors the eager
|
# Refresh the static write cursor in-place (mirrors the eager
|
||||||
# snapshot-then-advance). Skip the advance during capture: dummy slots
|
# snapshot-then-advance). Skip the advance during capture: dummy slots
|
||||||
# would corrupt real ring positions.
|
# would corrupt real ring positions.
|
||||||
|
|||||||
@@ -0,0 +1,63 @@
|
|||||||
|
"""Regression test: Mamba2 (NemotronH) ``extra_buffer`` state tracking stays
|
||||||
|
consistent under CUDA-graph decode.
|
||||||
|
|
||||||
|
The decode-cache-hit KL check guards the graph-replay track-save path: a decode
|
||||||
|
that crosses ``--mamba-track-interval`` donates its tracked state to the mamba
|
||||||
|
radix tree, the second turn hits that decode-seeded prefix, and its logprobs
|
||||||
|
must match a cold recompute. Before the ``[:bs]`` track-buffer fix in
|
||||||
|
``MambaAttnBackendBase._replay_metadata``, the captured track-save kernel bound
|
||||||
|
the stale tail of the static track buffer and scattered the state to the wrong
|
||||||
|
slot, so the tree cached an unwritten slot and this check fails with KL ~1.5
|
||||||
|
(vs ~1e-3 healthy).
|
||||||
|
|
||||||
|
The existing ``extra_buffer`` KL coverage (test_unified_radix_cache_kl_mamba,
|
||||||
|
the int8 checkpoint e2e) runs GDN models (Qwen3-Next), whose decode track-save
|
||||||
|
indexes the track buffer from the front and never hit the bug; NemotronH
|
||||||
|
(Mamba2) slices it from the tail (``[-num_decodes:]``), so it needs its own
|
||||||
|
registered coverage.
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
python3 -m unittest test_mamba2_extra_buffer_kl
|
||||||
|
"""
|
||||||
|
|
||||||
|
import unittest
|
||||||
|
|
||||||
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
from sglang.test.kits.kl_divergence_kit import KLDivergenceMixin
|
||||||
|
from sglang.test.server_fixtures.default_fixture import DefaultServerBase
|
||||||
|
|
||||||
|
register_cuda_ci(est_time=600, stage="extra-a", runner_config="1-gpu-large")
|
||||||
|
|
||||||
|
|
||||||
|
class TestMamba2ExtraBufferKL(KLDivergenceMixin, DefaultServerBase):
|
||||||
|
"""NemotronH (Mamba2) + extra_buffer: cache-hit logprobs match cold recompute."""
|
||||||
|
|
||||||
|
model = "nvidia/NVIDIA-Nemotron-Nano-9B-v2"
|
||||||
|
|
||||||
|
# Decode-seeded reuse is the regression trigger (the graphed decode
|
||||||
|
# track-save); the broken path fails at KL ~1.5, so 0.005 discriminates
|
||||||
|
# cleanly while absorbing bf16 reuse noise. Prefill reuse (chunk-aligned
|
||||||
|
# intermediate h states) is inherently looser; threshold matches the manual
|
||||||
|
# TestNvidiaNemotronNanoV2BF16ExtraBuffer calibration.
|
||||||
|
kl_div_thres = 0.005
|
||||||
|
kl_div_thres_prefill = 0.01
|
||||||
|
kl_div_max_samples = 16
|
||||||
|
|
||||||
|
other_args = [
|
||||||
|
"--max-mamba-cache-size",
|
||||||
|
"256",
|
||||||
|
"--mem-fraction-static",
|
||||||
|
"0.8",
|
||||||
|
"--mamba-scheduler-strategy",
|
||||||
|
"extra_buffer",
|
||||||
|
# The 512-token decode turns must cross a track boundary for the tree
|
||||||
|
# to hold decode-seeded states; halve the default interval (must stay
|
||||||
|
# >= the model's mamba chunk size, 128) so nearly every sample donates
|
||||||
|
# one even when it stops early.
|
||||||
|
"--mamba-track-interval",
|
||||||
|
"128",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user