Skip to content

Fix minifloat RMSNorm output dtype, add E3M4 MXFP format, drop state_dict deepcopy - #322

Open
Shreyas8612 wants to merge 4 commits into
mainfrom
fix/rmsnorm-dtype-and-mxfp-e3m4
Open

Shreyas8612 wants to merge 4 commits into
mainfrom
fix/rmsnorm-dtype-and-mxfp-e3m4

Conversation

@Shreyas8612

Copy link
Copy Markdown
Collaborator

Three small, independent fixes on the quantised-module path, each with a
regression test.

  • LlamaRMSNormMinifloat and Qwen3RMSNormMinifloat returned float32 for
    bf16 input whenever the replacement module's weight was still float32
    (float32 * bf16 promotes the product). The product is now cast back to
    the input dtype, matching the upstream RMSNorm contract, so the
    following linear sees the model dtype.
  • MXFPMeta accepts E3M4 (exp=3, frac=4) alongside E4M3/E5M2; the
    minifloat element quantiser already handled it, only the legality check
    rejected it.
  • weight_replacement() no longer deep-copies the source state_dict()
    before load_state_dict(), which copies each tensor into the target
    anyway; this removes a transient second copy of every replaced layer.
    The unused import is dropped.

Tests: test/nn/quantized/modules/test_rms_norm_minifloat_dtype.py,
test/nn/quantizers/test_mxfp_meta_formats.py and
test/passes/module/test_weight_replacement.py (18 tests, CPU only). The
dtype and format tests fail on main and pass with this branch;
test/passes/module still passes; black reports no changes.

Shreyas8612 added 4 commits September 15, 2026 20:59
weight_replacement() deep-copied the source module's full state_dict
before handing it to load_state_dict(). load_state_dict() already
copies each tensor into the target module's own parameters, so the
deepcopy only doubled peak host memory for the module being replaced,
which is significant when swapping large decoder layers or MLP experts
on a memory-limited host.

Pass the state_dict straight through and drop the now-unused deepcopy
import. Behaviour is unchanged: the source module is not mutated.

Verified by compiling the module and running the quantized-module
tests.
MXFPMeta rejected element_exp_bits=3, element_frac_bits=4, so an 8-bit
E3M4 MXFP configuration could not be built even though the underlying
minifloat quantiser handles it (it only requires exp + frac < 16).
E3M4 is a standard 8-bit MX element format alongside E4M3 and E5M2 and
is useful for weight formats that need more mantissa than E4M3.

Add (3, 4) to legal_element_exp_frac_bits and lay the tuple out one
entry per line, grouped by element width, so future additions are
obvious in review.

Verified by constructing MXFPMeta(block_size=32, scale_exp_bits=8,
element_exp_bits=3, element_frac_bits=4, ...) and quantising a tensor
with it.
LlamaRMSNormMinifloat and Qwen3RMSNormMinifloat returned
`weight * hidden_states.to(input_dtype)`. When the replacement module
is constructed outside the model's bf16 context its weight parameter
is float32, and float32 * bf16 promotes the result to float32, so a
bf16 model receives float32 hidden states from every norm. The next
linear then fails (float32 activations against bf16 weights) or, with
quantised linears, silently runs the activation quantiser on a wider
dtype than intended.

Cast the whole product to the input dtype so the module always returns
what it was given, matching the upstream RMSNorm contract.

Verified by instantiating both modules with a float32 weight, feeding a
bf16 tensor and checking the output is bf16 (it was float32 before the
change), and by running the quantized-module tests.
…cement

The previous three commits had no committed coverage. Add small CPU tests
that pin each contract:

- LlamaRMSNormMinifloat and Qwen3RMSNormMinifloat return the input dtype
  when the module weight is still float32 (quantised and bypass configs),
  and float32 input stays float32 with reference numerics.
- MXFPMeta constructs and quantises with every supported element format,
  including E3M4, and still rejects unsupported ones.
- weight_replacement gives the target its own copy and leaves the source
  untouched when the target is mutated afterwards.

The dtype and format tests fail against the previous rms_norm.py and
meta.py; all 18 pass with the fixes.
Copilot AI lite review requested due to automatic review settings September 17, 2026 01:46

Copilot AI 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.

🟡 Changes recommended

The new weight_replacement regression test uses .data and does not verify bias independence, so it should be adjusted before relying on it.

Get a fresh assessment by requesting another Copilot review.

Pull request overview

This PR applies three focused fixes in the quantized-module pathway: ensuring minifloat RMSNorm preserves input dtype, extending MXFP metadata to accept the E3M4 element format, and reducing memory overhead in weight_replacement() by removing an unnecessary state_dict() deep-copy, each accompanied by regression tests.

Changes:

  • Ensure LlamaRMSNormMinifloat and Qwen3RMSNormMinifloat return outputs in the input dtype even when weights remain fp32 after module replacement.
  • Allow MXFP element format E3M4 in MXFPMeta legality checks (in addition to existing formats).
  • Remove deepcopy() of state_dict() in weight_replacement() to avoid a transient second copy of parameters; add regression coverage.
File summaries
File Description
src/chop/nn/quantized/modules/llama/rms_norm.py Cast RMSNorm output product back to input dtype to prevent bf16→fp32 promotion.
src/chop/nn/quantized/modules/qwen3/rms_norm.py Same dtype preservation fix for Qwen3 minifloat RMSNorm.
src/chop/nn/quantizers/mxfp/meta.py Extend legal element exp/frac bit pairs to include E3M4.
src/chop/passes/module/module_modify_helper.py Drop deepcopy() from weight_replacement() to avoid duplicating tensors during loads.
test/nn/quantized/modules/test_rms_norm_minifloat_dtype.py Regression tests for dtype preservation behavior (bf16 and fp32 cases).
test/nn/quantizers/test_mxfp_meta_formats.py Tests validating accepted/rejected MXFP element formats, including E3M4.
test/passes/module/test_weight_replacement.py Regression test ensuring replacement copies are independent of the source module.
Review details
  • Files reviewed: 7/7 changed files
  • Comments generated: 2
  • Review effort level: Lite

💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.

Comment on lines +16 to +23
assert torch.equal(target.weight, source.weight)
assert torch.equal(target.bias, source.bias)
assert target.weight.data_ptr() != source.weight.data_ptr()

# Mutating the target must not leak back into the source.
target.weight.data.fill_(0.0)
assert torch.equal(source.weight, source_weight)
assert torch.equal(source.bias, source_bias)
Comment on lines +103 to 104
target_state_dict = x.state_dict()
missing_keys, unexpected_keys = y.load_state_dict(target_state_dict, strict=False)

This branch has not been deployed

No deployments
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.

2 participants