From b805cc501444a6b98e26e3088851ba5e980704d7 Mon Sep 17 00:00:00 2001 From: simple-sun <69040952+simple-sun@users.noreply.github.com> Date: Sat, 12 Sep 2026 09:38:09 +0800 Subject: [PATCH] [Bugfix] Track DFlash Mamba state at checkpoint boundaries (#37818) Co-authored-by: lvweiv Co-authored-by: kpham-sgl Co-authored-by: Claude Fable 5.1 Co-authored-by: Baizhou Zhang --- python/sglang/srt/speculative/dflash_worker_v2.py | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/python/sglang/srt/speculative/dflash_worker_v2.py b/python/sglang/srt/speculative/dflash_worker_v2.py index a95674399..391465d16 100644 --- a/python/sglang/srt/speculative/dflash_worker_v2.py +++ b/python/sglang/srt/speculative/dflash_worker_v2.py @@ -1777,6 +1777,7 @@ class DFlashWorkerV2(BaseSpecWorker): *, batch: ScheduleBatch, seq_lens_pre_verify: torch.Tensor, + seq_lens_post_verify: torch.Tensor, commit_lens: torch.Tensor, ) -> None: """Commit Mamba intermediate states for accepted verify steps. @@ -1796,10 +1797,10 @@ class DFlashWorkerV2(BaseSpecWorker): mamba_track_interval = mamba_track_grid(batch.tree_cache.page_size) to_track_mask = ( seq_lens_pre_verify // mamba_track_interval - != batch.seq_lens // mamba_track_interval + != seq_lens_post_verify // mamba_track_interval ) tracking_point = ( - batch.seq_lens // mamba_track_interval * mamba_track_interval + seq_lens_post_verify // mamba_track_interval * mamba_track_interval ) to_track_ith = torch.clamp(tracking_point - seq_lens_pre_verify - 1, min=0) can_track_mask = to_track_mask & ( @@ -2460,9 +2461,12 @@ class DFlashWorkerV2(BaseSpecWorker): if self._need_mamba_verify_commit: assert seq_lens_pre_verify is not None + if new_seq_lens is None: + new_seq_lens = prefix_lens + commit_lens.to(prefix_lens.dtype) self._update_target_mamba_state_after_verify( batch=batch, seq_lens_pre_verify=seq_lens_pre_verify, + seq_lens_post_verify=new_seq_lens, commit_lens=commit_lens, )