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:
sychen52
2026-09-29 10:31:31 -07:00
committed by GitHub
co-authored by Claude Opus 5.5
parent 94d7272e8b
commit 834c90d7a1
4 changed files with 136 additions and 65 deletions
@@ -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)