Skip to content

Fix Mamba state progression across cached prefill continuations - #7714

Open
wdykas wants to merge 1 commit into
NVIDIA:mainfrom
wdykas:fix/mamba-cached-continuation-state
Open

wdykas wants to merge 1 commit into
NVIDIA:mainfrom
wdykas:fix/mamba-cached-continuation-state

Conversation

@wdykas

@wdykas wdykas commented Sep 30, 2026 •

Copy link
Copy Markdown
Contributor
  • I, the PR author, have personally reviewed every line of this PR.

Problem

Hybrid models can skip Mamba recurrent state updates when a continuation prefill chunk matches cached attention KV blocks, producing predictions from a state that represents the wrong token history.

Attention KV blocks and Mamba snapshots have independent retention. If a Mamba snapshot is evicted while its KV blocks remain, the first prompt chunk correctly recomputes tokens until a usable state boundary. On subsequent block-aligned chunks, _compute_prefix_match falls through to the attention-only skip calculation: its hybrid fallback is guarded by finished == 0.

However, add_request restores Mamba state only on the first chunk. A continuation carries the live state from the previous chunk, so matching KV blocks alone cannot justify advancing past additional recurrent transitions. Restoring an earlier state on the first chunk does not make later skips safe.

Fix

Apply the hybrid zero-skip fallback to continuation chunks as well:

- elif self.is_hybrid_model and finished == 0:
+ elif self.is_hybrid_model:

First-chunk skipping still uses a matching Mamba snapshot. Matched attention KV blocks remain shared, and attention-only continuation skipping is unchanged. Hybrid continuations recompute the transitions necessary to keep the live recurrent state aligned with the logical prompt position.

Controlled reproduction

A four-GPU test on the affected deployment revision (d37db1077cb77ec08fb846485accd92ca39cd639) held model weights and the same 49,326-token prompt fixed. No weight update occurred between conditions.

Cache condition Missing recurrent transitions KL from cold next-token distribution
Cold 0 0
Warm, intact caches 0 0.0000128
Only Mamba snapshots removed 39,680 2.85843
Both caches cleared 0 0
Mamba snapshots removed, equivalent continuation guard applied 0 0

The failing trace skipped four 8,192-token chunks and one 6,912-token region without restoring Mamba state. The most likely next token changed: the cold top token's probability fell from 99.95% to 5.72%, while another token received 79.03%. The guard restored the entire cold next-token distribution exactly in this test.

This demonstrates an inference correctness defect. Its contribution to the RL reasoning-length divergence that prompted the investigation is still unproven.

Validation

  • 9 focused regression cases pass against both the deployment revision and this PR's current-main implementation. They cover absent snapshot storage, evicted snapshots, an earlier usable snapshot, later snapshots, both matcher call modes, and preserved attention-only skipping. They check that recurrent progress accounts for every prompt token while retaining matched KV blocks.
  • The focused tests use the real DynamicInferenceContext implementation with CPU tensors, through a one-process torch.distributed.run invocation; GPUs are hidden and unrelated distributed suite fixtures are omitted. The full GPU unit suite has not been run.
  • tools/autoformat.sh completed in check mode. Black, isort, Pylint, and Ruff passed. Its non-gating mypy step reported missing dependency stubs and type errors; this was not a clean mypy run.

Use fresh caches when adopting the fix: an earlier faulty request may already have stored inconsistent Mamba state or suffix KV.

Pre-checks

  • Added focused unit regression tests.
  • Documented the state invariant in the implementation.
  • Ran the repository formatter in check mode, with the mypy limitation above.
  • Added a functional test to the repository suite (the controlled GPU reproduction was run separately).

No linked issue. This PR does not add or change a GPU kernel or public API.

Signed-off-by: wdykas <wdykas@nvidia.com>
@copy-pr-bot

copy-pr-bot Bot commented Sep 30, 2026

Copy link
Copy Markdown

This pull request requires additional validation before any workflows can run on NVIDIA's runners.

Pull request vetters can view their responsibilities here.

Contributors can view more details about this message here.

@wdykas
wdykas marked this pull request as ready for review September 30, 2026 15:19
@wdykas
wdykas requested review from a team as code owners September 30, 2026 15:19
@wdykas

wdykas commented Sep 30, 2026

Copy link
Copy Markdown
Contributor Author

/ok to test ab7f0e1

@svcnvidia-nemo-ci svcnvidia-nemo-ci added the Final Review PR is in the "final review" stage label Sep 30, 2026

This branch was successfully deployed

2 active deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

complexity: low Final Review PR is in the "final review" stage

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants