Files
nvfp4-megamoe-kernel/tests/unit/test_nvfp4_1_1_quant.py
biondizzle 80b6b79f9e NVFP4-1.1: FP4 quantization primitives for CuTeDSL kernels
- fp8_e4m3_from_float32: manual FP8 E4M3 cast (bias=7, exp 0-15 valid,
  NaN guard for exp=15/mant=7, mantissa overflow handling)
- fp8_e4m3_to_float32: dequantize FP8 E4M3 bit pattern back to Float32
- half_step_to_e2m1_idx: E2M1 step mapping (0-12 → 0-7)
- quantize_e2m1_nibble: per-element E2M1 quantize + sign + pack
- Verified 0/500 trial failures against Python reference
- Key fixes discovered during validation:
  1. FP8 E4M3 bias is 7, NOT 8
  2. Exponent range is 0-15 (exp=15/mant=7 is NaN; others valid)
  3. Subnormal formula: val = m * 2^(-9) = m/512 (NOT m/1024)
  4. Round-to-nearest-even (not round-half-up) for half_step and mantissa
  5. Mantissa overflow (round to 8) must increment exponent
2026-05-28 03:39:55 +00:00

179 lines
6.2 KiB
Python

"""
NVFP4-1.1 Phase 1: Verify FP4 quantization math in CuTeDSL.
Tests the fp4_quant.py functions on B200. Compares CuTeDSL kernel output
with Python reference (quantize_activation_nvfp4).
The kernel takes 16 BF16 values + global_scale, quantizes to NVFP4,
and writes FP4 packed bytes + FP8 scale byte to output tensors.
Uses cute.arch.load for scalar GMEM reads (proven pattern from the codebase).
For writes, uses the output tensor's iterator + offset pattern.
"""
import torch
import cutlass
import cutlass.cute as cute
import cutlass.torch as cutlass_torch
import sys
import os
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "../.."))
from dsv4.ops.quantize import quantize_activation_nvfp4, SF_VEC_SIZE
from dsv4.kernels.gemm.fp4_quant import (
fp8_e4m3_from_float32_manual,
fp8_e4m3_to_float32,
half_step_to_e2m1_idx,
quantize_e2m1_nibble,
)
@cute.kernel
def fp4_quant_test_kernel(
input_bf16: cute.Tensor, # (16,) BF16 — 16 input values
out_data: cute.Tensor, # (10,) Int32 — [0..7] = FP4 packed bytes, [8] = SF byte, [9] = debug
gs_scalar: cute.Tensor, # (1,) Float32 — global scale
):
"""Quantize 16 BF16 values to NVFP4 using fp4_quant functions.
Single-thread kernel (only thread 0 does work).
Grid: (1, 1, 1), Block: (32, 1, 1)
"""
tidx, _, _ = cute.arch.thread_idx()
if tidx == cutlass.Int32(0):
# Load global scale
gs = cute.arch.load(gs_scalar.iterator, cutlass.Float32)
# Load 16 BF16 values, convert to FP32, normalize by global_scale
vals_f32 = [cutlass.Float32(0.0)] * 16
for i in cutlass.range(16, unroll=1):
bf16_val = cute.arch.load(
input_bf16.iterator + i * cutlass.Int32(2), # BF16 = 2 bytes
cutlass.BFloat16,
)
vals_f32[i] = bf16_val.to(cutlass.Float32) / gs
# ── Compute per-16-element amax ──
amax = cutlass.Float32(0.0)
for i in cutlass.range(16, unroll=1):
v = vals_f32[i]
a = cute.math.fmax(v, cutlass.Float32(0.0) - v) # abs
amax = cute.math.fmax(amax, a)
# ── Block scale = amax / 6 ──
bsf_f32 = amax / cutlass.Float32(6.0)
# Underflow: if amax < 6 * 2^-9, force scale = 0
underflow_threshold = cutlass.Float32(6.0 * (2.0 ** -9))
if amax < underflow_threshold:
bsf_f32 = cutlass.Float32(0.0)
# ── FP8 E4M3 cast ──
sf_bits = fp8_e4m3_from_float32_manual(bsf_f32)
# ── Dequantize FP8 scale (round-trip) ──
bs_dequant = fp8_e4m3_to_float32(sf_bits)
# ── Quantize each value to E2M1 and pack ──
for i in cutlass.range(8, unroll=1):
nibble0 = quantize_e2m1_nibble(vals_f32[2 * i], bs_dequant)
nibble1 = quantize_e2m1_nibble(vals_f32[2 * i + 1], bs_dequant)
packed = (nibble1 << cutlass.Int32(4)) | nibble0
# Write packed byte as Int32
cute.arch.store(out_data.iterator + i * cutlass.Int32(4), packed, cutlass.Int32)
# ── Write FP8 scale byte ──
cute.arch.store(out_data.iterator + cutlass.Int32(8) * cutlass.Int32(4), sf_bits, cutlass.Int32)
# ── Debug: write bsf_f32 and bs_dequant as float ──
# out_data[9] is unused — let's skip for simplicity
def run_test():
"""Run the FP4 quantization test."""
device = "cuda"
N = 16
# Generate test input
torch.manual_seed(42)
x_bf16 = torch.randn(1, N, dtype=torch.bfloat16, device=device)
# Compute global scale (matching quantize_activation_nvfp4)
x_f32 = x_bf16.float()
amax_val = x_f32.abs().max().item()
global_scale = max(amax_val / (6.0 * 448.0), 1e-8)
# Python reference
ref_fp4, ref_sf = quantize_activation_nvfp4(x_bf16, global_scale)
ref_fp4_bytes = ref_fp4.view(torch.uint8).reshape(-1).cpu()
ref_sf_bytes = ref_sf.view(torch.uint8).cpu()
print(f"Input BF16 (first 8): {x_bf16[0, :8].cpu()}")
print(f"Global scale: {global_scale:.8f}")
print(f"Ref FP4 bytes: {ref_fp4_bytes}")
print(f"Ref SF byte: {ref_sf_bytes}")
# Prepare output tensor
out_data = torch.zeros(10, dtype=torch.int32, device=device)
gs_tensor = torch.tensor([global_scale], dtype=torch.float32, device=device)
# Convert to CuTe tensors
def to_cute(t):
ct = cutlass_torch.from_dlpack(t)
return ct.mark_layout_dynamic(leading_dim=cutlass_torch.get_leading_dim(t))
x_flat = x_bf16.reshape(N).contiguous()
input_c = to_cute(x_flat)
out_c = to_cute(out_data)
gs_c = to_cute(gs_tensor)
# Compile and run
import cuda.bindings.driver as cuda
stream = cuda.CUstream(torch.cuda.current_stream().cuda_stream)
print("\nCompiling kernel (first run may take a minute)...")
compiled = cute.compile(
fp4_quant_test_kernel,
input_c, out_c, gs_c,
stream,
)
print("Compiled. Running...")
compiled(input_c, out_c, gs_c, stream)
torch.cuda.synchronize()
# Extract results
our_fp4 = out_data[:8].to(torch.uint8).cpu()
our_sf = out_data[8].to(torch.uint8).cpu().item()
print(f"\nOur FP4 bytes: {our_fp4}")
print(f"Our SF byte: {our_sf}")
# Compare
fp4_match = torch.equal(our_fp4, ref_fp4_bytes[:8])
sf_match = our_sf == ref_sf_bytes[0].item()
if fp4_match and sf_match:
print("\n✅ PASS: FP4 quantization matches Python reference!")
return True
else:
print(f"\n❌ FAIL: FP4 match={fp4_match}, SF match={sf_match}")
if not fp4_match:
for i in range(8):
o = our_fp4[i].item()
r = ref_fp4_bytes[i].item()
if o != r:
print(f" Byte {i}: ours=0x{o:02x}, ref=0x{r:02x}")
if not sf_match:
print(f" SF: ours=0x{our_sf:02x}, ref=0x{ref_sf_bytes[0].item():02x}")
return False
if __name__ == "__main__":
print("=" * 60)
print("NVFP4-1.1 Phase 1: FP4 Quantization Math Test")
print("Verifies fp4_quant.py functions match Python reference")
print("=" * 60)
success = run_test()
exit(0 if success else 1)