mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
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:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user