mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
Fix NVFP4 fake quant zeroing blocks with small scales (#2549)
### What does this PR do? Type of change: Bug fix The dynamic NVFP4 Triton kernel (`fp4_fake_quant_block`, used on compute >= 8.9) and the Conv3D implicit-GEMM CUDA kernels replaced any FP8 block scale below 1e-5 with 1.0, so every block whose max |x| was below ~6e-5 was zeroed. The static Triton kernel, the CUDA extension fallback and NVFP4 export have no such floor; the floor was only guarding division by zero. The Triton kernel had its own copy of the scale/round code instead of the shared `nvfp4_scalar_quant`. - Triton: use the shared `nvfp4_scalar_quant` (zero only on a zero block scale). - Conv3D CUDA (fused kernel and standalone `fp4_fake_quant`): same rule. - Both: a zero, inf or NaN global amax uses a unit block scale, like the CUDA extension (the conv kernels returned NaN for inf/NaN before). Blocks with scale >= 1e-5 are unchanged. - New tests: power-of-two scaling of input and global amax (2^-10, 2^-20) scales the output by the same factor (Triton, standalone conv FP4, fused conv3d); invalid global amax gives unit-scale rounding. The conv test's Python reference drops the floor. - B200: 505 passed / 31 skipped (`tests/gpu/torch/quantization` NVFP4/FP4 files) and 180 passed (conv implicit GEMM + attention P-QDQ). - Negative control on `main`: the small-input tests fail on both kernels; conv also fails for inf/NaN global amax. ### Before your PR is "*Ready for review*" - Is this change backward compatible?: ✅ (only blocks with scale < 1e-5 or an inf/NaN global amax change, from zeros/NaN to correct values) - If you copied code from any other sources or added a new PIP dependency, did you follow guidance in `CONTRIBUTING.md`: N/A - Did you write any new necessary tests?: ✅ - Did you update Changelog?: N/A - Did you get Claude approval on this PR?: ❌ <!-- This is an auto-generated comment: release notes by coderabbit.ai --> ## Summary by CodeRabbit * **Bug Fixes** * FP4 quantization now preserves proportional output when inputs and valid global scales are reduced together, including at small scales. * Zero, infinite, or NaN global scales use a safe fallback, preserving inputs already representable in FP4. * Small positive block scales are no longer discarded by an absolute scale threshold; subnormal scale handling is covered across quantization paths. * **Tests** * Added coverage for scale consistency across input types and block sizes, invalid global scales, and subnormal FP8 block scales. <!-- end of auto-generated comment: release notes by coderabbit.ai --> --------- Signed-off-by: Shiyang Chen <shiychen@nvidia.com> Co-authored-by: Claude Opus 5.5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 5.5
parent
94d7272e8b
commit
834c90d7a1
@@ -421,6 +421,29 @@ class TestConv3dFP4QuantBlockSizes:
|
||||
f"fp4_block_size={fp4_block_size}: mean diff {mean_diff:.6e} too high"
|
||||
)
|
||||
|
||||
@pytest.mark.parametrize("fp4_block_size", [16, 128])
|
||||
def test_quant_small_inputs_not_zeroed(self, cuda_conv3d, fp4_block_size):
|
||||
"""Scaling the activation and its amax by a power of two scales the output by the same factor.
|
||||
|
||||
Activation blocks with small magnitudes must not be flushed to zero by an absolute scale floor.
|
||||
"""
|
||||
torch.manual_seed(0)
|
||||
x = torch.randn(1, 16, 8, 8, 8, device="cuda", dtype=torch.float32)
|
||||
w = torch.randn(32, 16, 3, 3, 3, device="cuda", dtype=torch.float32)
|
||||
act_amax = x.abs().max().unsqueeze(0)
|
||||
kwargs = {
|
||||
"stride": (1, 1, 1),
|
||||
"padding": (1, 1, 1),
|
||||
"dilation": (1, 1, 1),
|
||||
"quant_act": True,
|
||||
"fp4_block_size": fp4_block_size,
|
||||
}
|
||||
reference = cuda_conv3d(x, w, act_amax=act_amax, **kwargs)
|
||||
assert reference.abs().max() > 0
|
||||
factor = 2.0**-20
|
||||
out = cuda_conv3d(x * factor, w, act_amax=act_amax * factor, **kwargs)
|
||||
assert torch.equal(out, reference * factor)
|
||||
|
||||
def test_smaller_block_less_error(self, cuda_conv3d):
|
||||
"""Smaller FP4 block sizes should generally produce lower quantization error.
|
||||
|
||||
@@ -607,13 +630,15 @@ def _py_fp4_fake_quant_ref(x_flat, global_amax, block_size):
|
||||
block = x_np[b * block_size : (b + 1) * block_size]
|
||||
block_max = float(max(abs(v) for v in block))
|
||||
|
||||
# Scale quantization
|
||||
scaled = block_max / (6.0 * global_scale)
|
||||
scaled = min(scaled, 448.0)
|
||||
quantized_scale = fp8_e4m3_roundtrip(scaled) * global_scale
|
||||
if quantized_scale < 1e-5:
|
||||
# Scale quantization; a zero, inf or NaN global scale gives a unit block scale
|
||||
if 0.0 < global_scale < math.inf:
|
||||
scaled = block_max / (6.0 * global_scale)
|
||||
scaled = min(scaled, 448.0)
|
||||
quantized_scale = fp8_e4m3_roundtrip(scaled) * global_scale
|
||||
else:
|
||||
quantized_scale = 1.0
|
||||
inv_scale = 1.0 / quantized_scale
|
||||
# Only a zero block scale zeroes the block
|
||||
inv_scale = 1.0 / quantized_scale if quantized_scale > 0.0 else 0.0
|
||||
|
||||
for i in range(block_size):
|
||||
val = block[i]
|
||||
@@ -722,6 +747,48 @@ class TestFP4FakeQuantScale:
|
||||
# Block 1 exact values should be close to E2M1 levels
|
||||
assert out[8:].abs().max() <= 6.0 + 1e-5
|
||||
|
||||
def test_small_inputs_not_zeroed(self, cuda_fp4):
|
||||
"""Scaling the input and global amax by a power of two scales the output by the same factor.
|
||||
|
||||
Blocks with small magnitudes must not be flushed to zero by an absolute scale floor.
|
||||
"""
|
||||
torch.manual_seed(0)
|
||||
inp = torch.randn(64 * 16, device="cuda") * 10
|
||||
amax = inp.abs().max().unsqueeze(0)
|
||||
reference = cuda_fp4(inp, amax, 16)
|
||||
assert reference.abs().max() > 0
|
||||
for factor in (2.0**-10, 2.0**-20):
|
||||
assert torch.equal(cuda_fp4(inp * factor, amax * factor, 16), reference * factor)
|
||||
|
||||
@pytest.mark.parametrize("global_amax", [0.0, float("inf"), float("nan")])
|
||||
def test_invalid_global_amax_uses_unit_scale(self, cuda_fp4, global_amax):
|
||||
"""A zero, inf or NaN global amax uses a unit block scale, like the modelopt CUDA extension."""
|
||||
torch.manual_seed(0)
|
||||
e2m1 = torch.tensor([0, 0.5, 1, 1.5, 2, 3, 4, 6], device="cuda")
|
||||
inp = e2m1[torch.randint(0, 8, (64 * 16,), device="cuda")]
|
||||
inp = inp * (torch.randint(0, 2, inp.shape, device="cuda") * 2 - 1)
|
||||
out = cuda_fp4(inp, torch.tensor([global_amax], device="cuda"), 16)
|
||||
assert torch.equal(out, inp)
|
||||
|
||||
@pytest.mark.parametrize("scale_steps", [1, 2, 3, 4, 5, 6, 7, 12])
|
||||
def test_subnormal_fp8_block_scale(self, cuda_fp4, scale_steps):
|
||||
"""Block scales in the FP8 E4M3 subnormal range (multiples of 2^-9 below 2^-6) are exact.
|
||||
|
||||
A global amax of 6 * 448 gives a global scale of 1, so a block whose max is 6 * s gets the
|
||||
block scale s and E2M1 values times s come back unchanged. 12 * 2^-9 is a normal control.
|
||||
"""
|
||||
scale = scale_steps * 2.0**-9
|
||||
e2m1 = torch.tensor([0, 0.5, 1, 1.5, 2, 3, 4, 6], device="cuda")
|
||||
sign = torch.tensor([1, -1, 1, -1, -1, 1, -1, 1], device="cuda")
|
||||
inp = torch.cat([e2m1 * sign, e2m1 * -sign]) * scale
|
||||
amax = torch.tensor([6.0 * 448.0], device="cuda")
|
||||
assert torch.equal(cuda_fp4(inp, amax, 16), inp)
|
||||
assert torch.equal(_py_fp4_fake_quant_ref(inp, amax, 16), inp)
|
||||
if _triton_fp4_available():
|
||||
from modelopt.torch.kernels.quantization.gemm import fp4_fake_quant_block
|
||||
|
||||
assert torch.equal(fp4_fake_quant_block(inp.view(1, 16), amax[0]).view(-1), inp)
|
||||
|
||||
|
||||
class TestFP4FakeQuantBlockSizes:
|
||||
"""Test different block sizes."""
|
||||
|
||||
@@ -354,3 +354,35 @@ class Testfp4:
|
||||
f"Mean abs diff: {(output_static - output_dynamic).abs().mean()}\n"
|
||||
f"Max relative diff: {((output_static - output_dynamic).abs() / (output_dynamic.abs() + 1e-8)).max()}"
|
||||
)
|
||||
|
||||
@pytest.mark.skipif(
|
||||
not hasattr(triton_kernel, "fp4_fake_quant_block"),
|
||||
reason="fp4_fake_quant_block requires compute >= 8.9",
|
||||
)
|
||||
@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16])
|
||||
def test_fp4_kernel_scales_with_small_inputs(self, dtype):
|
||||
"""Scaling the input and global amax by a power of two scales the output by the same factor.
|
||||
|
||||
Blocks with small magnitudes must not be flushed to zero by an absolute scale floor.
|
||||
"""
|
||||
torch.manual_seed(0)
|
||||
x = (torch.randn(8, 64, device="cuda") * 10).to(dtype)
|
||||
reference = triton_kernel.fp4_fake_quant_block(x, x.float().abs().amax())
|
||||
assert reference.abs().amax() > 0
|
||||
for factor in (2.0**-10, 2.0**-20):
|
||||
scaled = triton_kernel.fp4_fake_quant_block(x * factor, x.float().abs().amax() * factor)
|
||||
assert torch.equal(scaled, reference * factor)
|
||||
|
||||
@pytest.mark.skipif(
|
||||
not hasattr(triton_kernel, "fp4_fake_quant_block"),
|
||||
reason="fp4_fake_quant_block requires compute >= 8.9",
|
||||
)
|
||||
@pytest.mark.parametrize("global_amax", [0.0, float("inf"), float("nan")])
|
||||
def test_fp4_kernel_invalid_global_amax_uses_unit_scale(self, global_amax):
|
||||
"""A zero, inf or NaN global amax uses a unit block scale, like the CUDA extension."""
|
||||
torch.manual_seed(0)
|
||||
e2m1 = torch.tensor([0, 0.5, 1, 1.5, 2, 3, 4, 6], device="cuda")
|
||||
x = e2m1[torch.randint(0, 8, (8, 64), device="cuda")]
|
||||
x = x * (torch.randint(0, 2, x.shape, device="cuda") * 2 - 1)
|
||||
output = triton_kernel.fp4_fake_quant_block(x, torch.tensor(global_amax, device="cuda"))
|
||||
assert torch.equal(output, x)
|
||||
|
||||
Reference in New Issue
Block a user