Conversation
…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>
…t) for Plan-B handoff" This reverts commit 9be502d.
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>
Thank you, @jackneil 🙏Just verified your fix on 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:
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: |
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).
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)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 passedmask=Noneand attended it. Asymmetric behaviour ⇒ Path A ≠ Path B.CompressedKVCache.accumulate— the prefill compressor withoverlap=Trueproduces 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_chunkslot and, when overlap is enabled, callingcompressor(prev || cur)[:, 1:2]to recreate prefill semantics exactly.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 withcomp_visible = (q_pos + 1) >= (k_idx + 1) * ras the prepend mask.DeepseekV4Model.__call__—create_attention_maskwas called withoutwindow_size=, so re-prefill mask was causal-only whileRotatingKVCachephysically dropped keys outside the window. Fixed by passingwindow_size=self.args.sliding_window._prev_chunkis also propagated throughmerge/filter/extend/extract/nbytesfor batched-cache parity.Verification
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=4BEFORE this fix: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]— PASStest_deepseek_v4_decode_cache_argmax_parity[4|8|12]— PASStest_rotating_kv_cache_meta_state_round_trip_post_wrap— PASStests/test_models.py -k deepseek_v4— 8 PASS, 1 skipped (no regressions)tests/test_models.py -k cache— 5 PASSKnown follow-up (not in this PR)
CompressedKVCache.state/meta_stateonly proxylocal's arrays —_pool/_buf/_buf_count/_prev_chunkare not yet round-tripped throughsave_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