fix(qwen3.5-VL): correct image inference (interleaved M-RoPE, patch-embed, compressed-position decode + delta lifetime) - #80
Merged
Conversation
The existing VL image-chat e2e tests (`*_t0_capture`) only assert paged==flat
byte-identity, which passes even when the shared vision path produces garbage
features (the model describes the ocr.png financial table as "stone/fabric").
Add a correctness gate per file — `qwen3_5_moe_vl_reads_document_text` (MoE,
primary) and `qwen3_5_vl_reads_document_text` (dense, shared vision path). Each
sends ONE image+prompt turn at T=0 ("Transcribe the text in this document...")
with max_new_tokens=512 past the small thinking budget, then asserts the output
contains >=2 of the document keywords (reconciliation/bank/council/trunch/
october/2019/balance), printing the full model output on failure.
This is the TDD ground-truth gate for the qwen3.5 VL image fix; it is
`#[ignore]`-gated on MLX_TEST_QWEN35{,MOE}_VL_MODEL_PATH + MLX_TEST_VLM_IMAGE_PATH
(image defaults to examples/ocr.png) and never runs under a plain `cargo test`.
Reuses the existing cfg/user_msg/resolve_image_path helpers; no new infra.
Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
qwen3_5 / qwen3_5_moe image inference rotated Q/K with PaddleOCR-VL's sectioned (contiguous-chunk) multimodal-RoPE selector instead of qwen3_5's interleaved (stride-3) selector. For image tokens (where temporal != height != width) this assigns ~17% of frequencies to the wrong spatial axis, producing garbage visual features. Text tokens were unaffected because temporal == height == width makes the three cos/sin axis rows bit-identical, so any selector yields the same angles — which is exactly why text worked and images did not. Add `apply_multimodal_rotary_pos_emb_interleaved`, which builds the per-frequency axis selector matching mlx-vlm's `_interleaved_position_selector` (height: idx 1,4,7,… up to section[1]*3; width: idx 2,5,8,… up to section[2]*3; temporal otherwise), mirrors it across the doubled cos/sin, and gathers the selected t/h/w axis per frequency via `take_along_axis` before the same rotate_half + partial-rotary tail. Switch qwen3_5 attention's two M-RoPE call sites (forward + forward_paged) to it; this covers qwen3_5_moe too (shared `Qwen3_5Attention`). PaddleOCR-VL keeps the sectioned path (unchanged; its production forward uses the C++ sectioned kernel). Tests (no model): interleaved selection correctness against the hand-computed selector, and a text-invariance hard gate asserting the interleaved apply is bit-identical to the sectioned apply when t==h==w (proving the text path cannot regress). Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
The qwen3_5 shared vision patch-embed loader collapsed the 5D Conv3d weight [out, kD=2, kH, kW, in] by taking only temporal slice 0 and never loaded the Conv bias. The image processor duplicates the static frame across the temporal axis (mlx-vlm qwen3_vl Conv3d(bias=True), kD=2), so the effective 2D kernel is the SUM of the temporal slices plus the bias. Sum over the temporal axis (robust to kD != 2) and plumb the optional patch_embed.proj.bias through set_patch_embed / PatchEmbedding::new into the Conv2d bias (was hardcoded None). get_parameters now round-trips the bias. Shared path (qwen3_5 / qwen3_5_moe / Qwen3.6-35B-A3B), not forked. Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
`MxArray::addmm` called `mlx_array_addmm`, but the vendored `mlx::core::addmm` in this build returns only `alpha*(A@B)` and silently drops `beta*C`. That was invisible for the bias-free LM linears (Q/K/V/O/MLP and the bias-free MoE router gate all pass a zero C), but it corrupted every biased linear — most visibly the Qwen3.5-VL vision tower, whose qkv/proj/fc1/fc2 and merger projections all carry a bias. Dropping those biases produced semantically wrong image features (the model described the ocr.png financial table as "stone/fabric"). Compute the result explicitly (matmul, optional alpha scale, then add beta*C) so the C term is actually applied. Numeric bisection vs mlx-vlm on converted bf16 ornith: vision-tower image_embeds per-token cosine median 0.9996 after the fix (block0 rel-err 35x -> 0.0078). Add nn::linear unit tests proving addmm applies a [4]/[1,4]/[2,4] C and Linear::forward applies its bias. Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
An image run compresses its placeholder tokens into fewer M-RoPE positions, so image prefill rotates Q/K at the compressed positions and records rope_deltas = max_position + 1 - seq_len (negative). Decode and warm continuation must keep rotating at that compressed position (physical_slot + rope_deltas) while K/V is still written at the physical slot. The scalar-offset RoPE path rotated at the raw physical token count, ignoring rope_deltas, so post-image decode ran ~725 positions off and produced repetition garbage. Carry rope_deltas on VisionMerge, store it as cached_rope_deltas at every VLM prefill, and thread a rope_position_offset (physical position + cached_rope_deltas, cast u32->i32 before the negative add) through forward_paged, the decoder-layer paged forward, and the paged decode/prefill drivers for both qwen3_5 and qwen3_5_moe. get_rope_index takes the global max over the (t,h,w) axes for the delta. Text turns store no delta, so the offset equals the physical position and non-VLM paths stay byte-identical. Matches paddleocr_vl (cache_offset + rope_deltas) and mlx-vlm (base_offset + rope_delta). Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
…che hits The cross-turn M-RoPE delta an image prefill bakes into the shared `cached_rope_deltas` was reset only when `cached_prefix_len == 0`. But the paged turn planner also produces `cached_prefix_len > 0` with `continued_live_prefix == false` on a non-live prefix-cache hit: a later text request that merely shares a cached text prefix with an earlier image request (the model is shared across all sessions). There the old gate did not fire, so a stale negative delta survived and rotated unrelated text at `physical_slot + stale_delta` instead of the raw physical position. Only a live continuation (`continued_live_prefix`) re-attends the image's physically-resident compressed-position K/V, so only it needs the delta: image requests prefill with `skip_lookup` and never publish a text stream that collides with their expanded-placeholder blocks, so every non-live hit restores pure-text prefix blocks (delta 0). Factor the decision into `rope_delta_for_paged_turn(current, continued_live_prefix)` and wire all six paged gates (dense + MoE sync, stream, engine paths) through it. Add model-free lifecycle regression tests covering the live-continuation, cold-start, and stale-delta-on-text-hit cases. Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
… mlx bug The earlier addmm rewrite attributed the dropped `beta*C` to a bug in `mlx::core::addmm`. That premise was wrong: PyPI MLX's addmm applies `C` correctly and the FFI wrapper passes its args correctly. The real cause was a corrupt local metallib that miscompiled the fused GEMM kernels (the same bad build also miscompiled the NAX gemm). The explicit matmul+add form is kept as robustness against this project's documented non-deterministic metallib corruption, and the nn::linear C-application tests double as a build canary — but the code comment no longer claims an mlx source bug. Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
|
Important Review skippedAuto reviews are disabled on this repository. Please check the settings in the CodeRabbit UI or the ⚙️ Run configurationConfiguration used: Organization UI Review profile: CHILL Plan: Pro Run ID: You can disable this status message by setting the Use the checkbox below for a quick retry:
✨ Finishing Touches🧪 Generate unit tests (beta)
Thanks for using CodeRabbit! It's free for OSS, and your support helps us grow. If you like it, consider giving us a shout-out. Comment |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Summary
Fixes Qwen3.5-VL image inference for both the
qwen3_5(dense) andqwen3_5_moefamilies. Before this branch, an image turn ran end-to-end butproduced garbage descriptions — the model could not read text and hallucinated
unrelated subjects (e.g. described a bank-reconciliation table as "dark grey
stone / white fabric"). Text inference and
mlx convertwere unaffected; onlythe visual-feature path was wrong.
Found via
ornith-1.0-35b(a Qwen3.5-VL-MoE-35B-A3B post-train). Root-caused bynumeric bisection against the reference
mlx-vlmimplementation.Root causes & fixes (all verified vs
mlx-vlm)Interleaved multimodal RoPE for image tokens (
fb95c6d8)Qwen3.5-VL uses the interleaved (stride-3 per-freq) M-RoPE selector, not
the PaddleOCR sectioned/contiguous one. Text tokens (
t==h==w) areidentical either way — which is exactly why text always worked while images
(
t≠h≠w) got the wrong rotation axis on ~17% of frequencies, every attentionlayer, destroying the 2D positional structure.
Patch embedding: sum Conv3d temporal slices + load
patch_embed.proj.bias(6a889bdb)The patch embed previously used only the first temporal slice and hard-coded
the conv bias to
None.Decode rotates at the compressed M-RoPE position (
b5c4420e)Image prefill compresses ~754 placeholder tokens into ~29 M-RoPE positions, so
rope_deltas = max_position + 1 − seq_len(≈ −726). Decode must rotate thequery at
physical_slot + rope_deltaswhile K/V still writes at the physicalslot. Threaded a
rope_position_offset: i32throughforward_paged/decoder_layer/paged_forward(dense + MoE share oneforward_paged); the negative delta is carried onVisionMerge.rope_deltasand stored in
cached_rope_deltasat VLM prefill. The physical position iscast
u32 → i32before the negative add to avoid underflow.rope-delta lifetime gate (
f7ae00ca)Found by adversarial review of (3). The cross-turn delta was reset only when
cached_prefix_len == 0, but the paged-turn planner also yieldscached_prefix_len > 0withcontinued_live_prefix == falseon a non-liveprefix-cache hit. Because the model instance is shared across all sessions,
a stale negative delta from a prior image turn could then leak into an
unrelated text-only request that merely shares a cached text prefix —
rotating that text at
physical + stale_delta→ garbage.Factored the decision into
rope_delta_for_paged_turn(cached_rope_deltas, continued_live_prefix)(keep the delta iff it is a live image continuation,else clear) and wired it through all six paged reset gates (dense + MoE ×
sync / stream / engine). This is safe because image requests prefill with
skip_lookupand never publish a hashable text stream that collides withtheir expanded-placeholder blocks, so every non-live hit restores only
pure-text prefix blocks (delta 0); only a live continuation re-attends the
image's compressed-position K/V. Added three model-free lifecycle regression
tests.
e2e correctness gate (
ab044b88)qwen3_5_moe_vl_reads_document_textandqwen3_5_vl_image_chat.rs(dense) —#[ignore], env-gated (MLX_TEST_QWEN35MOE_VL_MODEL_PATH+MLX_TEST_VLM_IMAGE_PATH), asserting the model reads ≥2DOC_KEYWORDSfromexamples/ocr.png. These are the ground-truth gates for VL image inference.Note on the addmm commit (
12e89b3a+8e6d69d9)12e89b3areplaced the fusedmlx_array_addmmprimitive with an explicitmatmul + add, originally attributed to a bug inmlx::core::addmm. Thatpremise was wrong — PyPI MLX's
addmmapplies theCterm correctly(maxdiff 0) and the FFI wrapper passes its arguments correctly. The real cause
was a corrupt local metallib that miscompiled the fused GEMM kernels (the same
bad build also miscompiled the NAX gemm). The explicit form is kept as
robustness against this project's documented non-deterministic metallib
corruption — it is correctness-equivalent and vision/bias-only so the perf cost
is negligible — and the
nn::linearC-application tests double as a buildcanary.
8e6d69d9corrects the misleading comment so it no longer claims an mlxsource bug.
Validation
ornith-1.0-35bcheckpoint: the modelreads the document, matching all 7 keywords
(
reconciliation, bank, council, trunch, october, 2019, balance).nn::linearaddmm/bias tests,paged_forwardrope-offset + delta-lifetimetests, and the broader rope/m-rope/delta unit sweep (70 tests) pass.
cargo clippy -p mlx-core --all-targets -D warningsclean;cargo fmtclean.Test plan
🤖 Generated with Claude Code
Note
High Risk
Changes core attention RoPE, vision weights, and shared paged inference state across dense/MoE VL; incorrect delta lifetime or offset math would corrupt multi-turn or cached-prefix text, though extensive unit and env-gated e2e tests mitigate this.
Overview
Fixes Qwen3.5-VL image inference (dense and MoE) by aligning vision, RoPE, and paged decode with
mlx-vlm.Interleaved M-RoPE — Adds
apply_multimodal_rotary_pos_emb_interleavedand switches Qwen3.5 attention from the PaddleOCR sectioned apply so image tokens get stride-3 per-frequency axis selection; text-only paths stay unchanged via invariance tests.Vision tower —
addmmis implemented as explicitmatmul+ scaled add (avoids local metallib fused-GEMM corruption that dropped biases). Patch embed loads optional conv bias and collapses Conv3d weights by summing temporal slices instead of taking slice 0.Compressed M-RoPE decode —
VisionMergecarriesrope_deltas;get_rope_indexuses the global max over t/h/w axes. Paged prefill/decode/MTP threadcached_rope_deltasthroughpaged_rope_offsetandrope_position_offsetso rotation uses compressed positions while KV stays at physical slots.rope_delta_for_paged_turnkeeps the delta only oncontinued_live_prefixso stale image deltas do not leak into unrelated text prefix-cache hits.Tests — Unit tests for interleaved RoPE, rope offset/delta lifecycle, patch embed, linear bias; env-gated e2e
reads_document_textgates for dense and MoE VL.Reviewed by Cursor Bugbot for commit 8e6d69d. Bugbot is set up for automated code reviews on this repo. Configure here.