Sitelet https://github.com/mlx-node/mlx-node/pull/78
Skip to content

fix(convert): apply imatrix AWQ on VLM checkpoints + compensate GDN in_proj_a/b - #78

Merged
Brooooooklyn merged 1 commit into
mainfrom
fix/awq-imatrix-vlm-prefix-and-gdn-ab
Jun 26, 2026
Merged

Brooooooklyn merged 1 commit into
mainfrom
fix/awq-imatrix-vlm-prefix-and-gdn-ab

Conversation

@Brooooooklyn

@Brooooooklyn Brooooooklyn commented Jun 25, 2026 •

Copy link
Copy Markdown
Contributor

Summary

Converting qwen-agentworld-35b-a3b (a qwen3_5_moe
*ForConditionalGeneration, GatedDeltaNet hybrid) to unsloth + mxfp8 with an
imatrix surfaced two real bugs in apply_awq_prescaling that silently degrade
AWQ 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 by
gguf_name_to_hf from blk.N.*). But sanitized VLM checkpoints carry
language_model.model.layers.N.*, and apply_awq_prescaling auto-detected that
prefix and used it for both the weight ops and the importance lookups →
every importance.get missed → modified 0, with no warning (the missing-key
warning only fires on a partial match).

Fix: imatrix_lookup_key() strips the language_model. wrapper before the
importance lookup; weight ops keep the detected prefix. apply_awq_prescaling
now returns the modified count, and the caller emits a warn! when it is 0
despite 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_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_* projections read the same
input_layernorm output:

  • decoder_layer.rs:203,207 — normed = input_layernorm.forward(x) → gdn.forward(&normed, …)
  • gated_delta_net.rs:264/270-271 — that normed feeds both in_proj_qkvz and in_proj_ba (= b ++ a)

So leaving in_proj_a/in_proj_b unscaled divided their inputs by s (here
~0.18–5.5×/channel) with un-compensated weights → corrupted GDN decay (a) and
beta (b) gates. The reparametrization is output-preserving only if every
consumer of the divided norm is column-scaled by s.

Fix: column-scale in_proj_a/in_proj_b by the same s (their 8-bit-affine
bit-width is orthogonal and unchanged). Also corrects the stale doc comment that
claimed a/b "have no preceding norm" — only o_proj/out_proj lack 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.

modified == 0  before fix 1 (prefix mismatch → no-op)
modified == 7  before fix 2 (a/b skipped)
modified == 9  after both fixes

Validation

  • cargo test -p mlx-core --lib convert:: → 104 passed; clippy + fmt clean.
  • Real convert of qwen-agentworld-35b-a3b with --q-recipe unsloth --q-mxfp --q-bits 4 --imatrix-path … (67GB bf16 → ~20GB, 4 shards) → coherent
    generation.
  • On-artifact AWQ fingerprint conv/(orig+1): per-channel non-uniform on
    input_layernorm (Group C full-attn + Group D GDN), ≈1.0 on
    post_attention_layernorm (Group A correctly inert — no dense FFN in MoE).

Note: --q-recipe unsloth --q-mode mxfp8 is rejected by the CLI (recipes allow
only affine|nvfp4); the recipe+micro-scaling path is --q-mxfp.

🤖 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 strip language_model. via imatrix_lookup_key, matching canonical model.layers.* entries from GGUF imatrix files. Previously every lookup missed and AWQ did nothing with no signal. apply_awq_prescaling now 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_layernorm and column-scales in_proj_qkv / in_proj_z, it also column-scales in_proj_a and in_proj_b with 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_weights expecting modified == 9.

Reviewed by Cursor Bugbot for commit 43468d0. Bugbot is set up for automated code reviews on this repo. Configure here.

@coderabbitai

coderabbitai Bot commented Jun 25, 2026 •

Copy link
Copy Markdown

Important

Review skipped

Auto reviews are disabled on this repository. Please check the settings in the CodeRabbit UI or the .coderabbit.yaml file in this repository. To trigger a single review, invoke the @coderabbitai review command.

⚙️ Run configuration

Configuration used: Organization UI

Review profile: CHILL

Plan: Pro

Run ID: 6312076f-0624-481c-9f2c-5d416f03a0f3

You can disable this status message by setting the reviews.review_status to false in the CodeRabbit configuration file.

Use the checkbox below for a quick retry:

  • 🔍 Trigger review
✨ Finishing Touches
🧪 Generate unit tests (beta)
  • Create PR with unit tests
  • Commit unit tests in branch fix/awq-imatrix-vlm-prefix-and-gdn-ab

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.

❤️ Share

Comment @coderabbitai help to get the list of available commands.

…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
Brooooooklyn force-pushed the fix/awq-imatrix-vlm-prefix-and-gdn-ab branch from 323d53d to 43468d0 Compare June 25, 2026 16:28
@Brooooooklyn
Brooooooklyn merged commit 409a7b5 into main Jun 26, 2026
8 checks passed
@Brooooooklyn
Brooooooklyn deleted the fix/awq-imatrix-vlm-prefix-and-gdn-ab branch June 26, 2026 00:44
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>
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.

1 participant