From fa354a1733c159bb6e6a6ffefaff7c660d0b5552 Mon Sep 17 00:00:00 2001 From: Shreyas8612 Date: Tue, 15 Sep 2026 20:59:11 +0100 Subject: [PATCH 1/4] Avoid deep-copying the state dict during module replacement 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. --- src/chop/passes/module/module_modify_helper.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/src/chop/passes/module/module_modify_helper.py b/src/chop/passes/module/module_modify_helper.py index 441457d9b..c3574c7f4 100644 --- a/src/chop/passes/module/module_modify_helper.py +++ b/src/chop/passes/module/module_modify_helper.py @@ -1,7 +1,6 @@ import torch from functools import reduce, partial -from copy import deepcopy import logging import inspect @@ -101,7 +100,7 @@ def check_module_instance(module, prefix_map): def weight_replacement(x, y): - target_state_dict = deepcopy(x.state_dict()) + target_state_dict = x.state_dict() missing_keys, unexpected_keys = y.load_state_dict(target_state_dict, strict=False) if missing_keys: logging.warning( From c8fcb764bac52f502c7a6e2449debf82281030a7 Mon Sep 17 00:00:00 2001 From: Shreyas8612 Date: Tue, 15 Sep 2026 20:59:40 +0100 Subject: [PATCH 2/4] Add E3M4 to the legal MXFP element formats 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. --- src/chop/nn/quantizers/mxfp/meta.py | 12 +++++++++++- 1 file changed, 11 insertions(+), 1 deletion(-) diff --git a/src/chop/nn/quantizers/mxfp/meta.py b/src/chop/nn/quantizers/mxfp/meta.py index fcd66bf26..4f5aff6c7 100644 --- a/src/chop/nn/quantizers/mxfp/meta.py +++ b/src/chop/nn/quantizers/mxfp/meta.py @@ -38,7 +38,17 @@ def __post_init__(self): f"Invalid scale exponent bits: {self.scale_exp_bits}. " f"Legal values are: {legal_scale_exp_bits}." ) - legal_element_exp_frac_bits = ((4, 3), (5, 2), (2, 3), (3, 2), (2, 1), (1, 2)) + # (exp_bits, frac_bits) element formats in the accelerator search + # space: 4-bit E1M2/E2M1, 6-bit E2M3/E3M2, 8-bit E3M4/E4M3/E5M2. + legal_element_exp_frac_bits = ( + (1, 2), + (2, 1), + (2, 3), + (3, 2), + (3, 4), + (4, 3), + (5, 2), + ) el_exp_frac = (self.element_exp_bits, self.element_frac_bits) assert el_exp_frac in legal_element_exp_frac_bits, ( f"Invalid element exp/frac bits: {el_exp_frac}. " From fb0a1cc83708504566056dd85f88a6c919d840b5 Mon Sep 17 00:00:00 2001 From: Shreyas8612 Date: Tue, 15 Sep 2026 21:01:31 +0100 Subject: [PATCH 3/4] Keep minifloat RMSNorm output in the input dtype 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. --- src/chop/nn/quantized/modules/llama/rms_norm.py | 5 ++++- src/chop/nn/quantized/modules/qwen3/rms_norm.py | 5 ++++- 2 files changed, 8 insertions(+), 2 deletions(-) diff --git a/src/chop/nn/quantized/modules/llama/rms_norm.py b/src/chop/nn/quantized/modules/llama/rms_norm.py index f777ef8a7..37d98bbdb 100644 --- a/src/chop/nn/quantized/modules/llama/rms_norm.py +++ b/src/chop/nn/quantized/modules/llama/rms_norm.py @@ -81,4 +81,7 @@ def forward(self, hidden_states): if self.w_quantizer is not None else self.weight ) - return weight * hidden_states.to(input_dtype) + # Cast the product, not just hidden_states: after module replacement + # self.weight may still be float32 while the activations are bf16, and + # float32 * bf16 promotes the output back to float32. + return (weight * hidden_states.to(input_dtype)).to(input_dtype) diff --git a/src/chop/nn/quantized/modules/qwen3/rms_norm.py b/src/chop/nn/quantized/modules/qwen3/rms_norm.py index 9241f6b4d..933f33406 100644 --- a/src/chop/nn/quantized/modules/qwen3/rms_norm.py +++ b/src/chop/nn/quantized/modules/qwen3/rms_norm.py @@ -58,4 +58,7 @@ def forward(self, hidden_states): if self.w_quantizer is not None else self.weight ) - return weight * hidden_states.to(input_dtype) + # Cast the product, not just hidden_states: after module replacement + # self.weight may still be float32 while the activations are bf16, and + # float32 * bf16 promotes the output back to float32. + return (weight * hidden_states.to(input_dtype)).to(input_dtype) From 4fa7b3404637e3db4fe7a032a86a4914c60fb142 Mon Sep 17 00:00:00 2001 From: Shreyas8612 Date: Wed, 16 Sep 2026 19:21:46 +0100 Subject: [PATCH 4/4] Add regression tests for RMSNorm dtype, MXFP formats and weight replacement 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. --- .../modules/test_rms_norm_minifloat_dtype.py | 61 +++++++++++++++++++ test/nn/quantizers/test_mxfp_meta_formats.py | 42 +++++++++++++ test/passes/module/test_weight_replacement.py | 23 +++++++ 3 files changed, 126 insertions(+) create mode 100644 test/nn/quantized/modules/test_rms_norm_minifloat_dtype.py create mode 100644 test/nn/quantizers/test_mxfp_meta_formats.py create mode 100644 test/passes/module/test_weight_replacement.py diff --git a/test/nn/quantized/modules/test_rms_norm_minifloat_dtype.py b/test/nn/quantized/modules/test_rms_norm_minifloat_dtype.py new file mode 100644 index 000000000..bb8f39c64 --- /dev/null +++ b/test/nn/quantized/modules/test_rms_norm_minifloat_dtype.py @@ -0,0 +1,61 @@ +"""The minifloat RMSNorm modules must return tensors in the input dtype. + +Module replacement constructs the quantised RMSNorm before the model is moved +to its working dtype, so ``weight`` can still be float32 while the activations +are bfloat16. ``float32 * bfloat16`` promotes to float32, and the following +linear then fails on a dtype mismatch. The modules cast the product back to the +input dtype; these tests pin that contract. +""" + +import pytest +import torch +from transformers import LlamaConfig, Qwen3Config + +from chop.nn.quantized.modules.llama.rms_norm import LlamaRMSNormMinifloat +from chop.nn.quantized.modules.qwen3.rms_norm import Qwen3RMSNormMinifloat + +HIDDEN = 16 +EPS = 1e-6 + +QUANTISED = { + "weight_exponent_width": 4, + "weight_frac_width": 3, + "data_in_exponent_width": 4, + "data_in_frac_width": 3, +} +BYPASS = {"bypass": True} + +CASES = [ + (LlamaRMSNormMinifloat, LlamaConfig(hidden_size=HIDDEN, rms_norm_eps=EPS)), + (Qwen3RMSNormMinifloat, Qwen3Config(hidden_size=HIDDEN, rms_norm_eps=EPS)), +] + + +def _reference(weight: torch.Tensor, x: torch.Tensor) -> torch.Tensor: + x32 = x.float() + normed = x32 * torch.rsqrt(x32.pow(2).mean(-1, keepdim=True) + EPS) + return (weight * normed.to(x.dtype)).to(x.dtype) + + +@pytest.mark.parametrize("cls,config", CASES, ids=["llama", "qwen3"]) +@pytest.mark.parametrize("q_config", [QUANTISED, BYPASS], ids=["quantised", "bypass"]) +def test_bf16_input_with_fp32_weight_returns_bf16(cls, config, q_config): + module = cls(config=config, q_config=q_config) + assert module.weight.dtype == torch.float32 + + hidden = torch.randn(3, HIDDEN).to(torch.bfloat16) + out = module(hidden) + + assert out.dtype == torch.bfloat16 + assert out.shape == hidden.shape + if q_config is BYPASS: + torch.testing.assert_close(out, _reference(module.weight, hidden)) + + +@pytest.mark.parametrize("cls,config", CASES, ids=["llama", "qwen3"]) +def test_fp32_input_stays_fp32(cls, config): + module = cls(config=config, q_config=BYPASS) + hidden = torch.randn(3, HIDDEN) + out = module(hidden) + assert out.dtype == torch.float32 + torch.testing.assert_close(out, _reference(module.weight, hidden)) diff --git a/test/nn/quantizers/test_mxfp_meta_formats.py b/test/nn/quantizers/test_mxfp_meta_formats.py new file mode 100644 index 000000000..e2b615718 --- /dev/null +++ b/test/nn/quantizers/test_mxfp_meta_formats.py @@ -0,0 +1,42 @@ +"""MXFPMeta accepts every supported element format and rejects the rest.""" + +import pytest +import torch + +from chop.nn.quantizers.mxfp import mxfp_quantizer_sim +from chop.nn.quantizers.mxfp.meta import MXFPMeta + + +def _meta(exp_bits: int, frac_bits: int) -> MXFPMeta: + return MXFPMeta( + block_size=32, + scale_exp_bits=8, + element_exp_bits=exp_bits, + element_frac_bits=frac_bits, + element_is_finite=True, + round_mode="rn", + ) + + +@pytest.mark.parametrize( + "exp_bits,frac_bits", + [(1, 2), (2, 1), (2, 3), (3, 2), (3, 4), (4, 3), (5, 2)], + ids=["E1M2", "E2M1", "E2M3", "E3M2", "E3M4", "E4M3", "E5M2"], +) +def test_legal_element_formats_construct_and_quantise(exp_bits, frac_bits): + meta = _meta(exp_bits, frac_bits) + x = torch.randn(4, 64) + q = mxfp_quantizer_sim(x, block_dim=-1, mxfp_meta=meta) + assert q.shape == x.shape + assert torch.isfinite(q).all() + + +def test_e3m4_is_an_eight_bit_format(): + meta = _meta(3, 4) + assert meta.element_exp_bits + meta.element_frac_bits + 1 == 8 + + +@pytest.mark.parametrize("exp_bits,frac_bits", [(3, 5), (6, 1), (1, 1)]) +def test_unsupported_element_formats_are_rejected(exp_bits, frac_bits): + with pytest.raises(AssertionError): + _meta(exp_bits, frac_bits) diff --git a/test/passes/module/test_weight_replacement.py b/test/passes/module/test_weight_replacement.py new file mode 100644 index 000000000..d6d692f5a --- /dev/null +++ b/test/passes/module/test_weight_replacement.py @@ -0,0 +1,23 @@ +"""weight_replacement copies parameters into the target without touching the source.""" + +import torch + +from chop.passes.module.module_modify_helper import weight_replacement + + +def test_target_receives_an_independent_copy(): + source = torch.nn.Linear(4, 4) + target = torch.nn.Linear(4, 4) + source_weight = source.weight.detach().clone() + source_bias = source.bias.detach().clone() + + weight_replacement(source, target) + + 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)