Skip to content

fix(deepseek_v4): with-cache decode parity + regression tests (@anerjy) - #3

Open
jackneil wants to merge 6 commits into
machiabeli:feat/deepseek-v4from
jackneil:fix/deepseek-v4-decode-cache
Open

jackneil wants to merge 6 commits into
machiabeli:feat/deepseek-v4from
jackneil:fix/deepseek-v4-decode-cache

Conversation

@jackneil

@jackneil jackneil commented Apr 26, 2026 •

Copy link
Copy Markdown

Summary

Fixes the @anerjy decode-cache bug on ml-explore#1189 and ships the deterministic regression tests that prove it. Original bug: with-cache autoregressive decode produced different logits at S=1 than re-prefilling the same token sequence, so greedy generation visibly corrupted after the first token ('{"verdict":"yes"}' → '{"ver":"yes"}').

Investigation went one layer deeper than the original bisection (which had ruled out caches, rope, compressor, sinks, and mHC residual). Per-layer instrumentation exposed a chain of four mutually-reinforcing defects in the compressed-attention + sliding-window paths. Each fix unmasked the next.

Root causes (all in mlx_lm/models/deepseek_v4.py)

  1. Bool-dtype prepend mask was zero-False in V4Attention.__call__ — for a bool mask, mx.zeros(... dtype=bool) is False = "block", silently masking out the entire compressed pool from S>1 (prefill / re-prefill) attention while S=1 decode passed mask=None and attended it. Asymmetric behaviour ⇒ Path A ≠ Path B.
  2. Ratio-4 cross-chunk overlap not reproducible at decode emission in CompressedKVCache.accumulate — the prefill compressor with overlap=True produces row K (K≥1) using chunk K-1's first-half features + chunk K's second-half features (via _overlap_transform). Decode-time emission only saw the current chunk, missing the first-half. Fixed by carrying a _prev_chunk slot and, when overlap is enabled, calling compressor(prev || cur)[:, 1:2] to recreate prefill semantics exactly.
  3. Non-causal pool prepend mask in V4Attention.__call__ — once Bug 1 was inverted to "include compressed pool", non-last queries in re-prefill saw pool rows summarising future tokens, while Path B's incrementally-grown pool only contained already-emittable rows. Fixed with comp_visible = (q_pos + 1) >= (k_idx + 1) * r as the prepend mask.
  4. Sliding-window cutoff missing on no-cache forward in DeepseekV4Model.__call__ — create_attention_mask was called without window_size=, so re-prefill mask was causal-only while RotatingKVCache physically dropped keys outside the window. Fixed by passing window_size=self.args.sliding_window.

_prev_chunk is also propagated through merge / filter / extend / extract / nbytes for batched-cache parity.

Verification

pytest tests/test_deepseek_v4_decode_cache.py -v --timeout=300
  7 passed  (3 logit-parity at atol=1e-5, 3 argmax-parity, meta-state round-trip)

pytest tests/test_models.py --timeout=300
  83 passed, 2 skipped  (no regressions in V4 / cache / other models)

The original bisection log (.bisect-results/v4-bisect-results.md) and the regression tests remain in-repo as permanent guards. The xfail markers from the handoff revision are removed (the tests are now hard guards).

Diagnostic from the smallest failing case (kept here for the record)

prompt_len=4 BEFORE this fix:

no-cache argmax:   [461, 326, 726, 459, 702, 102, 676, 979]
with-cache argmax: [461, 326,  10, 102, 461,  65, 449, 102]
                          ^-- divergence at index 2 (cache.offset 5→6)

After the fix, both paths produce byte-identical token sequences across all parameterized prompt lengths (4 / 8 / 12) and 8 decode steps each.

Test plan

  • test_deepseek_v4_decode_cache_logit_parity[4|8|12] — PASS
  • test_deepseek_v4_decode_cache_argmax_parity[4|8|12] — PASS
  • test_rotating_kv_cache_meta_state_round_trip_post_wrap — PASS
  • tests/test_models.py -k deepseek_v4 — 8 PASS, 1 skipped (no regressions)
  • tests/test_models.py -k cache — 5 PASS

Known follow-up (not in this PR)

CompressedKVCache.state / meta_state only proxy local's arrays — _pool / _buf / _buf_count / _prev_chunk are not yet round-tripped through save_prompt_cache / load_prompt_cache. Pre-existing gap, orthogonal to this fix; deserves its own commit + serialization round-trip test.

Reported-by: @anerjy
Co-Authored-By: Claude Opus 4.7 (1M context) noreply@anthropic.com

Jack Neil and others added 5 commits April 26, 2026 11:51
…erjy bug)

Adds a deterministic regression test for the decode-cache bug @anerjy
reported on mlx-lm PR ml-explore#1189 (ml-explore#1189).

The test drives a fixed token sequence through two paths on a tiny
synthetic DeepSeek-V4 model:
  Path A: re-prefill the full sequence at every step (no cache, ground truth)
  Path B: standard autoregressive decode with make_prompt_cache

It asserts the last-position logits (and argmax token) match between paths
through 8 decode steps. On unpatched PR ml-explore#1189 HEAD, this test fails — the
with-cache path diverges from ground truth at S=1+, corrupting output past
the first generated token.

The fix for the underlying bug ships in a follow-up commit; this commit
ensures the regression test exists first, per TDD discipline.

Reported-by: anerjy <https://github.com/anerjy>
Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
…lan-B handoff

After 4 hours of bisection (8 toggles, all in .bisect-results/), root cause
is localized to V4Attention's core attention path (NOT in caches/rope-kernel/
compressors/sinks/mHC) but not isolated to a small patch. Wrapping the parity
tests in pytest.mark.xfail(strict=True) so:

- CI reports them as XFAIL (not failure) on the buggy upstream branch
- Once a maintainer fixes the bug, the tests XPASS, failing CI under strict
  mode and reminding the fixer to drop the xfail decorator and convert the
  xfail into a hard regression guard

The meta_state round-trip test (which passes today) is intentionally NOT
marked xfail.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
Path A (no-cache, re-prefill ground truth) and Path B (with-cache,
autoregressive single-step decode) were diverging from decode step 2 onward
on prompt_len=4 and progressively earlier on longer prompts. Root cause was
not a single bug but a chain of four mutually-reinforcing defects in the
compressed-attention + sliding-window code paths. Bisecting them required
per-layer instrumentation; each fix exposed the next layer of the onion.

Bug 1 — bool-dtype prepend mask was zero-False (V4Attention __call__):
    comp_mask = mx.zeros(comp_shape, dtype=mask.dtype)
For a bool mask, zeros == False == "block", silently masking out the entire
compressed pool from attention in the S>1 (prefill / re-prefill) path while
the S=1 decode path passes mask=None (compressed pool fully attended). Fixed
by branching on dtype and additionally applying causal pool masking (Bug 3).

Bug 2 — ratio-4 overlap not reproducible at decode emission
(CompressedKVCache.accumulate): the prefill compressor with overlap=True
produces row K (K>=1) that mixes chunk K-1's first-half features with chunk
K's second-half features (via _overlap_transform). The decode-time emission
called compressor(_buf[:r]), which only sees chunk K's tokens — equivalent
to a 1-chunk compressor call, missing the cross-chunk first-half. Fixed by
storing the most-recent fully-consumed chunk as _prev_chunk and, when
emitting at decode under overlap=True, calling compressor(prev || cur)[:,
1:2] to recreate prefill semantics exactly. Also propagated _prev_chunk
through merge / filter / extend / extract / nbytes for batched-cache paths.

Bug 3 — non-causal pool prepend (V4Attention __call__): once Bug 1 was
inverted to "include compressed pool", the prepend marked the pool as
unconditionally visible to every query position. But pool row K summarizes
tokens up to position (K+1)*r - 1, so a query at position p < (K+1)*r - 1
would attend a row containing FUTURE tokens in the re-prefill forward.
Path B's incrementally-grown pool naturally only contained
already-emittable rows, so the two paths diverged at every non-last query
position and any downstream non-compressed layer cached the divergent K/V.
Fixed by computing comp_visible = (q_pos + 1) >= (k_idx + 1) * r and using
that as the prepend mask (with proper bool / additive-float branches).

Bug 4 — sliding-window cutoff missing on no-cache forward
(DeepseekV4Model __call__): create_attention_mask was called WITHOUT
window_size=, so the no-cache mask is causal-only while RotatingKVCache
physically drops keys outside the sliding window. Once offset >
sliding_window, Path A's last-position attention sees more keys than Path
B has. Fixed by passing window_size=self.args.sliding_window.

Verification:
  pytest tests/test_deepseek_v4_decode_cache.py -v --timeout=300
    7 passed (3 logit-parity at atol=1e-5, 3 argmax-parity, meta-state)
  pytest tests/test_models.py --timeout=300
    83 passed, 2 skipped (no regressions in V4 / cache / other models)

The CompressedKVCache state pickle (server prompt-cache reuse) does not
yet serialize _pool / _buf / _buf_count / _prev_chunk; that pre-existing
gap is orthogonal to this fix and would deserve its own commit + test.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
@jackneil jackneil changed the title test(deepseek_v4): regression tests for decode-cache bug (@anerjy) fix(deepseek_v4): with-cache decode parity + regression tests (@anerjy) Apr 26, 2026
@anerjy

anerjy commented Apr 26, 2026

Copy link
Copy Markdown

Thank you, @jackneil 🙏

Just verified your fix on mlx-community/DeepSeek-V4-Flash-bf16 against my original reproducer:

NO-CACHE  : '{"verrid":"yes"}<eos>...'
WITH-CACHE: '{"verrid":"yes"}<eos>...'
MATCH: True ✅

Byte-for-byte parity between cache-on and cache-off paths — exactly what we needed.

Going one diagnostic layer deeper than my (admittedly incomplete) bisection paid off. I had ruled out the surface culprits and concluded the divergence had to be ULP-level numeric drift amplified by mHC + Sinkhorn. Forcing fp32 SDPA + fp32 Sinkhorn didn't move the needle — that's when I tapped out and posted the "this might not be a single fixable bug" follow-up.

You found four mutually-reinforcing defects I didn't see, and the chain of unmasking is genuinely beautiful:

  1. Bool-dtype prepend mask zero=False — the prefill path silently masking out the entire compressed pool while decode (mask=None) attended it. So Path A and Path B were structurally inequivalent before any numerics even entered the picture. This is the kind of bug that hides in plain sight — mx.zeros(..., dtype=bool) reading as "block" is exactly the inverted convention I'd never have looked for.
  2. Ratio-4 cross-chunk overlap unreproducible at decode emission — the prefill compressor's _overlap_transform carries chunk K-1's first-half features into row K, but decode emission only saw the current chunk. The _prev_chunk carry is the right shape of fix and clearly the result of careful reading of _overlap_transform.
  3. Non-causal pool prepend mask — once Fix DeepSeek V4 checkpoint loading and mixed quantization #1 was inverted, the next layer of bug surfaced: non-last queries in re-prefill could see pool rows summarising future tokens. comp_visible = (q_pos + 1) >= (k_idx + 1) * r is exactly the right rule. Layered fixes that each unmask the next is the hardest debugging mode.
  4. Sliding-window cutoff missing on no-cache forward — RotatingKVCache physically dropping keys outside the window while create_attention_mask was causal-only. This one's especially sneaky because the cache's physical behavior and the mask's causality were each individually correct.

The 7 regression tests + the bisection log staying in-repo as permanent guards is the right call — these failure modes are exactly the kind that an unrelated future refactor would re-break silently.

Heads-up on the batch path for whichever maintainer picks this up: mlx_lm.server's batch generator passes an mx.array offset (from BatchRotatingKVCache) into V4Attention.__call__, and the q_pos = mx.arange(offset, offset + S).reshape(S, 1) line at the new causal-pool mask path crashes with TypeError: arange() incompatible argument types. Single-request flow (the one I tested in the repro) is unaffected. I'm working around it locally by hiding cache.merge via a property descriptor so the server falls back to the non-batch path, but that's a band-aid. Probably wants a if isinstance(offset, mx.array): offset = offset.max().item() (or the proper batched broadcast) in the same hunk that introduced q_pos.

vllm-mlx's BatchedEngine CompressedKVCache returns cache.offset as a
0-D mx.array rather than Python int. The mx.arange(offset, offset+S)
call at line 977 rejected mx.array positional args with TypeError.

Mirrors the existing coercion at RoPE __call__ (line 241-243) so the
two code paths handle cache.offset uniformly.

Verified live: DeepSeek-V4-Flash-2bit-DQ inference "What is 2+2?" returns
"4" with finish_reason=stop in 1.15s after this patch (vs every decode
step crashing with TypeError before).
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants