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) 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}. " 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( 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)