fix(convert): apply imatrix AWQ on VLM checkpoints + compensate GDN in_proj_a/b - #78
Merged
Merged
Conversation
|
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 |
…n_proj_a/b `apply_awq_prescaling` had two bugs that silently degraded unsloth conversions of VLM-wrapped (qwen3_5_moe `*ForConditionalGeneration`) and GatedDeltaNet models: 1. imatrix no-op on VLM checkpoints. imatrix importance is keyed canonical `model.layers.N.*` (from `gguf_name_to_hf`), but sanitized VLM weights carry `language_model.model.layers.N.*`, which AWQ auto-detected and used for the importance lookups too -> every lookup missed -> `modified 0`, no warning. Fix: normalize lookups via `imatrix_lookup_key()` (strip the `language_model.` wrapper); weight ops keep the detected prefix. `apply_awq_prescaling` now returns the modified count and the caller warns when it is 0 despite an imatrix being supplied. 2. AWQ Group D distorted GDN `in_proj_a`/`in_proj_b`. Group D divides the shared `input_layernorm` by per-channel `s` but scaled only `in_proj_qkv`/`in_proj_z`. All four `in_proj_*` read the same `input_layernorm` output (decoder_layer.rs:203,207 -> gated_delta_net.rs:264/270-271), so leaving a/b unscaled divided their inputs by `s` with un-compensated weights -> corrupted decay (`a`) / beta (`b`) gates. Fix: column-scale `in_proj_a`/`in_proj_b` by the same `s` (their 8-bit-affine bit-width is orthogonal/unchanged). Also correct the stale doc comment claiming a/b "have no preceding norm". Adds unit test `awq_prescaling_matches_vlm_prefixed_weights`: VLM-prefixed GDN + full-attention layers with a canonically-keyed imatrix must report `modified == 9` (was 0 before fix 1, 7 before fix 2). Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
Brooooooklyn
force-pushed
the
fix/awq-imatrix-vlm-prefix-and-gdn-ab
branch
from
June 25, 2026 16:28
323d53d to
43468d0
Compare
Brooooooklyn
added a commit
that referenced
this pull request
Jun 26, 2026
…79) ## Problem Main CI intermittently fails on a single Rust unit test (most recently on the PR #78 merge commit, run `28209793148`): ``` FAILED: nn::losses_test::tests::test_cross_entropy_qwen3_vocab crates/mlx-core/src/nn/losses_test.rs:630 "Loss should be > 10.0, got 9.881153" (1868 passed, 1 failed) ``` It is **unrelated to whatever PR happens to be merging** — the test was last touched in #12 and only `convert.rs` changed in #78. It is a pre-existing flaky test. ## Root cause The test builds random logits/targets with **no seed** and asserts a tight band on the mean-reduced loss: ```rust logits = random_normal([2, 151936], 0,1, None) // None is the DTYPE arg, not a seed targets = randint([2], 0, 151936) loss = mean_over_batch( logsumexp(logits) − logits[target] ) // ≈ log(151936) ≈ 11.93 assert!(10.0 < loss < 15.0) ``` With `batch_size = 2`, the loss is `≈ 12.43 − N(0, 0.71)` — a high random target logit occasionally drags the mean below `10.0` (the observed `9.88` is a ~3.5σ draw). ## Fix Bump `batch_size` 2 → 64. The mean-reduction then concentrates the loss at `log(V)`, shrinking the spread from ±0.71 to ±0.13. Vocab is unchanged, so the large-vocab **chunking path is still exercised**. ### Verification (2,000,000-draw simulation of the exact loss model) | batch | loss | min | P(loss<10) | P(loss>15) | |------:|------|----:|-----------:|-----------:| | 2 | 12.43 ± 0.71 | 9.01 | **0.030%** | 0.015% | | 64 | 12.43 ± 0.13 | 11.82 | **0** | 0 | At batch=64 the loss stays in `[11.8, 13.0]` across 2M draws — ~15σ from the `10.0` bound. (`logsumexp` constant measured = 12.4313, matching `ln(V)+0.5`.) 🤖 Generated with [Claude Code](https://claude.com/claude-code) <!-- CURSOR_SUMMARY --> --- > [!NOTE] > **Low Risk** > Test-only change with no production or loss-implementation edits; slightly more work per test run from a larger batch. > > **Overview** > Stabilizes **`test_cross_entropy_qwen3_vocab`** in `losses_test.rs` by raising **`batch_size` from 2 to 64** while keeping the Qwen3-sized vocab (151936) and the same `(10.0, 15.0)` loss band. > > The test still uses unseeded random logits/targets; with a tiny batch, a lucky target logit could drag the mean cross-entropy below 10.0 and fail CI. A larger batch tightens the mean-reduced loss around `log(vocab)` so the assertion is reliable without changing what the test exercises (large-vocab / chunked cross-entropy path). > > <sup>Reviewed by [Cursor Bugbot](https://cursor.com/bugbot) for commit 0225f73. Bugbot is set up for automated code reviews on this repo. Configure [here](https://www.cursor.com/dashboard/bugbot).</sup> <!-- /CURSOR_SUMMARY --> Co-authored-by: Claude Opus 4.8 <noreply@anthropic.com>
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
Converting
qwen-agentworld-35b-a3b(aqwen3_5_moe*ForConditionalGeneration, GatedDeltaNet hybrid) to unsloth + mxfp8 with animatrix surfaced two real bugs in
apply_awq_prescalingthat silently degradeAWQ for VLM-wrapped and GDN checkpoints. Both are fixed here with a unit
test, and validated end-to-end on the real model.
Bug 1 — imatrix AWQ was a silent no-op on VLM checkpoints
imatrix importance is always keyed canonical
model.layers.N.*(produced bygguf_name_to_hffromblk.N.*). But sanitized VLM checkpoints carrylanguage_model.model.layers.N.*, andapply_awq_prescalingauto-detected thatprefix and used it for both the weight ops and the importance lookups →
every
importance.getmissed →modified 0, with no warning (the missing-keywarning only fires on a partial match).
Fix:
imatrix_lookup_key()strips thelanguage_model.wrapper before theimportance lookup; weight ops keep the detected prefix.
apply_awq_prescalingnow returns the modified count, and the caller emits a
warn!when it is 0despite an imatrix being supplied (so this class of silent no-op is visible).
Bug 2 — AWQ Group D distorted GDN
in_proj_a/in_proj_bGroup D divides the shared
input_layernormby per-channelsbut scaled onlyin_proj_qkv+in_proj_z. All fourin_proj_*projections read the sameinput_layernormoutput:decoder_layer.rs:203,207—normed = input_layernorm.forward(x)→gdn.forward(&normed, …)gated_delta_net.rs:264/270-271— thatnormedfeeds bothin_proj_qkvzandin_proj_ba(= b ++ a)So leaving
in_proj_a/in_proj_bunscaled divided their inputs bys(here~0.18–5.5×/channel) with un-compensated weights → corrupted GDN decay (
a) andbeta (
b) gates. The reparametrization is output-preserving only if everyconsumer of the divided norm is column-scaled by
s.Fix: column-scale
in_proj_a/in_proj_bby the sames(their 8-bit-affinebit-width is orthogonal and unchanged). Also corrects the stale doc comment that
claimed a/b "have no preceding norm" — only
o_proj/out_projlack one.Test
awq_prescaling_matches_vlm_prefixed_weights: VLM-prefixed GDN (Group D) +full-attention (Group C) layers with a canonically-keyed imatrix must report
modified == 9.Validation
cargo test -p mlx-core --lib convert::→ 104 passed; clippy + fmt clean.qwen-agentworld-35b-a3bwith--q-recipe unsloth --q-mxfp --q-bits 4 --imatrix-path …(67GB bf16 → ~20GB, 4 shards) → coherentgeneration.
conv/(orig+1): per-channel non-uniform oninput_layernorm(Group C full-attn + Group D GDN), ≈1.0 onpost_attention_layernorm(Group A correctly inert — no dense FFN in MoE).🤖 Generated with Claude Code
Note
Medium Risk
Changes quantization weight math on convert for VLM and GDN models; incorrect scaling would affect model quality, but scope is limited to the AWQ path with an imatrix and is covered by a targeted regression test.
Overview
Fixes AWQ pre-scaling during convert so imatrix-based importance actually applies on VLM-wrapped checkpoints and stays mathematically consistent on GatedDeltaNet layers.
VLM / imatrix key mismatch: Weight keys stay on the detected prefix (e.g.
language_model.model.layers.*) while imatrix lookups now striplanguage_model.viaimatrix_lookup_key, matching canonicalmodel.layers.*entries from GGUF imatrix files. Previously every lookup missed and AWQ did nothing with no signal.apply_awq_prescalingnow returns a modified tensor count; the convert path warns when an imatrix was supplied but zero weights were touched.GDN Group D: When AWQ divides shared
input_layernormand column-scalesin_proj_qkv/in_proj_z, it also column-scalesin_proj_aandin_proj_bwith the same scale (importance still derived from qkv+z). Those projections share the normed input; skipping them left gates wrongly scaled. Docs for the unsloth recipe are updated to reflect this.Adds regression test
awq_prescaling_matches_vlm_prefixed_weightsexpectingmodified == 9.Reviewed by Cursor Bugbot for commit 43468d0. Bugbot is set up for automated code reviews on this repo. Configure here.