Fix numerical stability of test_gemm_common.py (#283)

Signed-off-by: Chenjie Luo <108829653+cjluo-nv@users.noreply.github.com>
This commit is contained in:
Chenjie Luo
2025-09-05 09:03:17 -07:00
committed by GitHub
parent 76fb12d476
commit 2b52759bcd
@@ -29,6 +29,12 @@ from modelopt.torch.quantization.backends.utils import fp4_compatible, fp8_compa
set_seed()
@pytest.fixture(autouse=True)
def setup_seed():
"""Set seed before each test function."""
set_seed()
@pytest.mark.parametrize(
("config", "gemm_forward", "atol", "rtol"),
[
@@ -257,9 +263,9 @@ def test_dynamic_gemm(model, config, gemm_forward, atol, rtol):
# The way the compression of the weights and inputs might be different.
# E.g. we may use torch.compile in the gemms.
assert torch.allclose(output_dynamic_quant_gemm, output_dynamic_quant, atol=atol / 3)
assert torch.allclose(output_calib_quant_gemm, output_calib_quant, atol=atol / 3)
assert torch.allclose(output_dynamic_quant_gemm, output_dynamic_quant, atol=atol / 2)
assert torch.allclose(output_calib_quant_gemm, output_calib_quant, atol=atol / 2)
assert torch.allclose(
output_dynamic_quant_gemm, output_dynamic_quant_compressed, atol=atol / 3
output_dynamic_quant_gemm, output_dynamic_quant_compressed, atol=atol / 2
)
assert torch.allclose(output_calib_quant_gemm, output_calib_quant_compressed, atol=atol / 3)
assert torch.allclose(output_calib_quant_gemm, output_calib_quant_compressed, atol=atol / 2)