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

feat(gemma4): load tied lm_head as packed-quantized (8-bit affine) - #77

Merged
Brooooooklyn merged 3 commits into
mainfrom
experiment/gemma4-lmhead8
Jun 25, 2026
Merged

Brooooooklyn merged 3 commits into
mainfrom
experiment/gemma4-lmhead8

Conversation

@Brooooooklyn

@Brooooooklyn Brooooooklyn commented Jun 25, 2026 •

Copy link
Copy Markdown
Contributor

What

Loads a quantized tied embedding / lm_head for gemma4 as a packed table instead of dequantizing the whole thing to dense bf16.

Before: a quantized embed_tokens went through Embedding::load_quantized, which dequantizes the entire table into a dense bf16 buffer and transposes it into embed_weight_t for the tied lm_head matmul (~2 GB bf16 for the 12B 262144x3840 vocab).

After: both tied and untied embeddings go through load_quantized_packed. forward() dequantizes only the gathered rows, and the tied lm_head projects through Embedding::as_linear (mlx_quantized_matmul, transpose=true) without ever materializing the dense table. The five logits sites in model.rs detect the packed backend via is_packed_quantized() and take the as_linear branch, so embed_weight_t stays None on the packed path.

The embedding quant mode is resolved from the per-layer config (affine or mxfp8) and threaded to the packed backend; biases-presence is driven off the actual embed_tokens.biases tensor key (affine has biases; mxfp8 does not). Any other quant mode at this key is rejected loud rather than silently mis-dequantizing.

Why / measured (gemma4 12B QAT-q4_0, M5 Max)

Quantizing the tied head from bf16 → 8-bit affine (gs32):

bf16 head 8-bit affine head
decode tok/s ~46 ~50 (equal within noise)
peak RSS 8.9 GB 8.1 GB (−0.8 GB)
on-disk 8.3 GB 7.5 GB (−0.8 GB)
greedy output — near-lossless (answers unchanged)

Validated coherent + near-lossless on both MLX 0.31.2 and 0.32.0. We also bench'd mxfp8 for this head — it's perf-tied with affine and only ~90 MB smaller, so 8-bit affine is the chosen default; the loader stays mode-generic so an mxfp8 tied head also loads.

Scope

Load-side support. It activates only when a checkpoint actually carries a quantized embed_tokens (.scales present + per-layer config). Dense bf16 tied models are unaffected (is_packed_quantized() is false → existing embed_weight_t path, unchanged). Auto-emitting a quantized tied head from mlx convert for gemma4 is a follow-up.

Validation

  • cargo clippy --all-targets -- -D warnings + cargo fmt --check: clean
  • Packed-embedding unit tests green, incl. packed_affine_as_linear_matches_dense_matmul (guards the exact packed-as_linear ≡ dense-matmul equivalence the tied lm_head relies on) and packed_mxfp8_forward_matches_dequant_full_then_gather
  • /codex:adversarial-review: approve, no material findings (verified transpose orientation, softcap-after-logits at all 5 sites, dense-tied + untied routing intact, no uncovered MTP/draft/warmup reader of embed_weight_t)

🤖 Generated with Claude Code


Note

Medium Risk
Changes core inference logits routing for quantized tied checkpoints; load-time guards reduce mis-dequant risk, but correctness depends on packed matmul matching the prior dense path.

Overview
Gemma4 no longer fully dequantizes a quantized embed_tokens into dense bf16 for a tied lm_head. Load now always uses load_quantized_packed (tied and untied), with resolve_packed_embed_params aligning affine vs mxfp8 to on-disk scales/biases and rejecting mis-resolved mxfp4/nvfp4-as-mxfp8 layouts.

At five logits sites in model.rs, when there is no separate lm_head and embeddings are packed, logits come from embed_tokens.as_linear (mlx_quantized_matmul) instead of embed_weight_t or a dense transpose of the full table—avoiding a large resident bf16 vocab matrix on the hot path.

Dense bf16 embeddings and an explicit lm_head are unchanged; the new behavior applies only when .scales are present and the packed backend is active.

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

@cursor cursor Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Cursor Bugbot has reviewed your changes and found 1 potential issue.

Fix All in Cursor

❌ Bugbot Autofix is OFF. To automatically fix reported issues with cloud agents, enable autofix in the Cursor dashboard.

Reviewed by Cursor Bugbot for commit b5bbb41. Configure here.

Comment thread crates/mlx-core/src/models/gemma4/persistence.rs Outdated
@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: 5e984a52-a63b-49d8-bc38-eef62287ceaf

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 experiment/gemma4-lmhead8

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.

@chatgpt-codex-connector chatgpt-codex-connector Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

💡 Codex Review

Here are some automated review suggestions for this pull request.

Reviewed commit: b5bbb41a7e

ℹ️ About Codex in GitHub

Your team has set up Codex to review pull requests in this repo. Reviews are triggered when you

  • Open a pull request for review
  • Mark a draft as ready
  • Comment "@codex review".

If Codex has suggestions, it will comment; otherwise it will react with 👍.

Codex can also answer questions or update the PR. Try commenting "@codex address that feedback".

Comment thread crates/mlx-core/src/models/gemma4/persistence.rs Outdated
@Brooooooklyn Brooooooklyn added the model-e2e Run the heavy Model E2E workflow (per-family real-checkpoint tests) on this PR label Jun 25, 2026
Brooooooklyn added a commit that referenced this pull request Jun 25, 2026
…cy for packed embed

Addresses PR #77 review (cursor + codex). The packed embed / tied-lm_head load
threaded the resolved PLQ's raw bits/group_size and a mode-only decision into
load_quantized_packed. Two hazards (latent on the shipped checkpoints, which
carry consistent explicit overrides, but reachable for override-less ones):

- mxfp8 without an explicit per-layer bits/group_size override falls back to
  `default_plq` = the affine body defaults (e.g. 4/64). MLX's quantized_matmul
  honors the passed bits/group_size rather than re-deriving them, so 4/64
  mis-unpacks the E8M0 table.
- the config mode could disagree with the on-disk tensors, letting an affine
  table (with .biases) be dequantized as mxfp8.

Add resolve_packed_embed_params, which makes (mode, bits, group_size, biases)
mutually consistent from the unambiguous tensor evidence before the packed
backend runs:
- mxfp8  -> force (MXFP8_GROUP_SIZE=32, MXFP8_BITS=8, "mxfp8"), biases=None
- affine -> bits/group_size from the PLQ, pass .biases through
- contradictions (mxfp8 + biases or non-uint8 scales; affine + uint8 scales)
  are rejected loud rather than producing garbage logits

Mirrors lfm2's plq_to_packed_params / try_build_mxfp8_quantized_linear, which
likewise force the MX constants and null biases for mxfp8. 4 unit tests added.

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
@Brooooooklyn

Copy link
Copy Markdown
Contributor Author

Addressed in ff651bda. Both comments verified against the control flow first — they're real but latent-only (the two shipped checkpoints carry consistent explicit embed_tokens overrides {bits:8, group_size:32, mode}, so neither fires on them; the hazard is reachable for an override-less checkpoint where embed_plq falls back to default_plq = the affine body defaults).

Root cause: the packed embed/tied-lm_head load threaded the resolved PLQ's raw bits/group_size + a mode-only decision into load_quantized_packed. Since MLX's quantized_matmul/fp_dequantize honor the passed bits/group_size, an mxfp8 table loaded with the affine 4/64 default would mis-unpack, and a config mode disagreeing with the on-disk tensors could dequantize an affine table (with .biases) as mxfp8.

Fix: new resolve_packed_embed_params makes (mode, bits, group_size, biases) mutually consistent from the unambiguous tensor evidence before the packed backend runs:

  • mxfp8 → force the MX pack constants (group_size=32, bits=8, mode="mxfp8") + biases=None — addresses @codex (P2).
  • affine → bits/group_size from the PLQ, pass .biases through.
  • contradictions fail loud (mxfp8 + .biases or non-uint8 scales; affine + uint8 scales) instead of producing garbage logits — addresses @cursor (the mode/biases disagreement).

Mirrors lfm2's plq_to_packed_params / try_build_mxfp8_quantized_linear, which already force the MX constants + null biases for mxfp8. The affine arm only rejects uint8 scales, which legitimate affine/sym8 embeddings never carry (their scales are bf16), so it can't over-reject existing checkpoints. 4 unit tests added; affine-Q8 + mxfp8 models both still decode coherently.

@chatgpt-codex-connector

Copy link
Copy Markdown

To use Codex here, create an environment for this repo.

@chatgpt-codex-connector chatgpt-codex-connector Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

💡 Codex Review

Here are some automated review suggestions for this pull request.

Reviewed commit: ff651bdad6

ℹ️ About Codex in GitHub

Your team has set up Codex to review pull requests in this repo. Reviews are triggered when you

  • Open a pull request for review
  • Mark a draft as ready
  • Comment "@codex review".

If Codex has suggestions, it will comment; otherwise it will react with 👍.

Codex can also answer questions or update the PR. Try commenting "@codex address that feedback".

Comment thread crates/mlx-core/src/models/gemma4/persistence.rs
Brooooooklyn and others added 3 commits June 25, 2026 20:12
…dequant

The tied embedding/lm_head load path routed quantized embeddings through
`load_quantized`, which dequantizes the whole table into a dense bf16 buffer
and transposes it into `embed_weight_t` for the lm_head matmul (~2 GiB bf16 for
the 12B 262144x3840 vocab). Switch to `load_quantized_packed` for both tied and
untied embeddings: `forward()` dequantizes only the gathered rows, and the tied
lm_head projects through `as_linear` (`mlx_quantized_matmul`) without ever
materializing the dense table. The five logits sites in model.rs detect the
packed backend via `is_packed_quantized()` and take the `as_linear` branch, so
`embed_weight_t` stays None on the packed path.

The quant mode is resolved from the per-layer config (affine or mxfp8) and
threaded to the packed backend; biases-presence is driven off the actual
`embed_tokens.biases` tensor key, not the mode alone. Any other quant mode at
this key is rejected loud rather than silently mis-dequantizing.

Measured (gemma4 12B, M5 Max): an 8-bit affine tied head saves ~0.8 GB RSS and
on-disk vs the bf16 tied head, with equal decode speed and near-lossless greedy
output.

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
…cy for packed embed

Addresses PR #77 review (cursor + codex). The packed embed / tied-lm_head load
threaded the resolved PLQ's raw bits/group_size and a mode-only decision into
load_quantized_packed. Two hazards (latent on the shipped checkpoints, which
carry consistent explicit overrides, but reachable for override-less ones):

- mxfp8 without an explicit per-layer bits/group_size override falls back to
  `default_plq` = the affine body defaults (e.g. 4/64). MLX's quantized_matmul
  honors the passed bits/group_size rather than re-deriving them, so 4/64
  mis-unpacks the E8M0 table.
- the config mode could disagree with the on-disk tensors, letting an affine
  table (with .biases) be dequantized as mxfp8.

Add resolve_packed_embed_params, which makes (mode, bits, group_size, biases)
mutually consistent from the unambiguous tensor evidence before the packed
backend runs:
- mxfp8  -> force (MXFP8_GROUP_SIZE=32, MXFP8_BITS=8, "mxfp8"), biases=None
- affine -> bits/group_size from the PLQ, pass .biases through
- contradictions (mxfp8 + biases or non-uint8 scales; affine + uint8 scales)
  are rejected loud rather than producing garbage logits

Mirrors lfm2's plq_to_packed_params / try_build_mxfp8_quantized_linear, which
likewise force the MX constants and null biases for mxfp8. 4 unit tests added.

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
…d at load)

Addresses PR #77 review (codex P2). is_mxfp8_checkpoint (hence
resolve_default_mode) classifies ANY uint8-scale checkpoint as mxfp8, but mxfp4
and nvfp4 also ship uint8 scales with no biases. For a mode-less/stale config a
mxfp4/nvfp4 embed_tokens can resolve to mode Mxfp8, reach the mxfp8 arm, and get
the 8-bit/group_size-32 constants forced onto it.

This does not silently mis-load (MLX's validate_quantized_input throws a shape
mismatch on the wrong constants), but the error is cryptic and deep in MLX.
Verify genuine mxfp8 packing at the loader instead: a real mxfp8 table has
weight_cols*4 == scales_groups*32 (both == hidden); mxfp4/nvfp4 halve the packed
weight width, so the check fails and we reject loud with a clear message naming
the mxfp4/nvfp4-vs-mxfp8 mismatch. Genuine mxfp8 and all affine paths unchanged.
1 unit test added.

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
@Brooooooklyn
Brooooooklyn force-pushed the experiment/gemma4-lmhead8 branch from ff651bd to 9b34999 Compare June 25, 2026 12:13
@Brooooooklyn

Copy link
Copy Markdown
Contributor Author

Synced onto latest main (rebased on #76) and addressed the new comment in 9b34999d.

@codex P2 (reject non-MXFP8 uint8-scale embeddings) — verified against the control flow first. The premise is real: is_mxfp8_checkpoint (and so resolve_default_mode) classifies any uint8-scale checkpoint as mxfp8, but mxfp4/nvfp4 also use uint8 scales + no biases — so a mode-less/stale config could resolve a mxfp4/nvfp4 embed_tokens to Mxfp8 and reach the mxfp8 arm.

One correction to the framing, though: this never silently mis-loads. MLX's validate_quantized_input throws a shape mismatch when the forced 8/32 constants don't match the packed table (and pre-ff651bda it threw a bits-mismatch), so it already failed loud — just with a cryptic error deep in MLX.

Fix: the loader now confirms genuine mxfp8 packing before forcing the constants — a real mxfp8 table has weight_cols*4 == scales_groups*32 (both == hidden); mxfp4/nvfp4 halve the packed weight width, so the check fails and we reject at the loader with a clear message naming the mxfp4/nvfp4-vs-mxfp8 mismatch. Genuine mxfp8 (the …-lmhead-mxfp8 model still decodes "Paris") and all affine paths are unchanged. Unit test added (resolve_packed_embed_mxfp4_shapes_resolved_to_mxfp8_fails_loud).

@chatgpt-codex-connector

Copy link
Copy Markdown

To use Codex here, create an environment for this repo.

@Brooooooklyn
Brooooooklyn merged commit 4a474a9 into main Jun 25, 2026
19 of 20 checks passed
@Brooooooklyn
Brooooooklyn deleted the experiment/gemma4-lmhead8 branch June 25, 2026 16:26
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

model-e2e Run the heavy Model E2E workflow (per-family real-checkpoint tests) on this PR

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant