Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 4 additions & 1 deletion src/chop/nn/quantized/modules/llama/rms_norm.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
5 changes: 4 additions & 1 deletion src/chop/nn/quantized/modules/qwen3/rms_norm.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
12 changes: 11 additions & 1 deletion src/chop/nn/quantizers/mxfp/meta.py
Original file line number Diff line number Diff line change
Expand Up @@ -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}. "
Expand Down
3 changes: 1 addition & 2 deletions src/chop/passes/module/module_modify_helper.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,6 @@
import torch

from functools import reduce, partial
from copy import deepcopy
import logging
import inspect

Expand Down Expand Up @@ -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)
Comment on lines +103 to 104
if missing_keys:
logging.warning(
Expand Down
61 changes: 61 additions & 0 deletions test/nn/quantized/modules/test_rms_norm_minifloat_dtype.py
Original file line number Diff line number Diff line change
@@ -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))
42 changes: 42 additions & 0 deletions test/nn/quantizers/test_mxfp_meta_formats.py
Original file line number Diff line number Diff line change
@@ -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)
23 changes: 23 additions & 0 deletions test/passes/module/test_weight_replacement.py
Original file line number Diff line number Diff line change
@@ -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)
Comment on lines +16 to +23
Loading