feat(gemma4): load tied lm_head as packed-quantized (8-bit affine) - #77
Conversation
There was a problem hiding this comment.
Cursor Bugbot has reviewed your changes and found 1 potential issue.
❌ 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.
|
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 |
There was a problem hiding this comment.
💡 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".
…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>
|
Addressed in Root cause: the packed embed/tied-lm_head load threaded the resolved PLQ's raw Fix: new
Mirrors lfm2's |
|
To use Codex here, create an environment for this repo. |
There was a problem hiding this comment.
💡 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".
…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>
ff651bd to
9b34999
Compare
|
Synced onto latest @codex P2 (reject non-MXFP8 uint8-scale embeddings) — verified against the control flow first. The premise is real: One correction to the framing, though: this never silently mis-loads. MLX's Fix: the loader now confirms genuine mxfp8 packing before forcing the constants — a real mxfp8 table has |
|
To use Codex here, create an environment for this repo. |

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_tokenswent throughEmbedding::load_quantized, which dequantizes the entire table into a dense bf16 buffer and transposes it intoembed_weight_tfor the tied lm_head matmul (~2 GB bf16 for the 12B262144x3840vocab).After: both tied and untied embeddings go through
load_quantized_packed.forward()dequantizes only the gathered rows, and the tied lm_head projects throughEmbedding::as_linear(mlx_quantized_matmul,transpose=true) without ever materializing the dense table. The five logits sites inmodel.rsdetect the packed backend viais_packed_quantized()and take theas_linearbranch, soembed_weight_tstaysNoneon the packed path.The embedding quant mode is resolved from the per-layer config (
affineormxfp8) and threaded to the packed backend; biases-presence is driven off the actualembed_tokens.biasestensor 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):
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(.scalespresent + per-layer config). Dense bf16 tied models are unaffected (is_packed_quantized()is false → existingembed_weight_tpath, unchanged). Auto-emitting a quantized tied head frommlx convertfor gemma4 is a follow-up.Validation
cargo clippy --all-targets -- -D warnings+cargo fmt --check: cleanpacked_affine_as_linear_matches_dense_matmul(guards the exact packed-as_linear≡ dense-matmul equivalence the tied lm_head relies on) andpacked_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 ofembed_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_tokensinto dense bf16 for a tied lm_head. Load now always usesload_quantized_packed(tied and untied), withresolve_packed_embed_paramsaligning 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 separatelm_headand embeddings are packed, logits come fromembed_tokens.as_linear(mlx_quantized_matmul) instead ofembed_weight_tor 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_headare unchanged; the new behavior applies only when.scalesare 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.