Fix minifloat RMSNorm output dtype, add E3M4 MXFP format, drop state_dict deepcopy - #322
Open
Shreyas8612 wants to merge 4 commits into
Open
Shreyas8612 wants to merge 4 commits into
Shreyas8612 wants to merge 4 commits into
Conversation
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.
There was a problem hiding this comment.
🟡 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
LlamaRMSNormMinifloatandQwen3RMSNormMinifloatreturn outputs in the input dtype even when weights remain fp32 after module replacement. - Allow MXFP element format E3M4 in
MXFPMetalegality checks (in addition to existing formats). - Remove
deepcopy()ofstate_dict()inweight_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
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.
Three small, independent fixes on the quantised-module path, each with a
regression test.
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.
minifloat element quantiser already handled it, only the legality check
rejected it.
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.