One-pass serialized NVFP4 weight expansion for BF16 consumer-GPU compute.
Functions:
fastvideo.layers.quantization.nvfp4_dequant.dequantize_nvfp4_cuda
dequantize_nvfp4_cuda(packed: Tensor, scales: Tensor, global_scale: float, dtype: dtype = bfloat16) -> Tensor
Expand E2M1 nibbles and swizzled E4M3 scales without full FP32 intermediates.
Source code in fastvideo/layers/quantization/nvfp4_dequant.py
| def dequantize_nvfp4_cuda(packed: torch.Tensor,
scales: torch.Tensor,
global_scale: float,
dtype: torch.dtype = torch.bfloat16) -> torch.Tensor:
"""Expand E2M1 nibbles and swizzled E4M3 scales without full FP32 intermediates."""
if not packed.is_cuda or scales.device != packed.device:
raise ValueError("NVFP4 fused dequantization requires tensors on the same CUDA device")
if packed.ndim != 2 or packed.dtype != torch.uint8 or scales.dtype != torch.uint8:
raise ValueError("NVFP4 fused dequantization requires packed uint8 weights and scales")
if not packed.is_contiguous() or not scales.is_contiguous():
raise ValueError("NVFP4 fused dequantization requires contiguous tensors")
rows, cols = packed.shape[0], packed.shape[1] * 2
if rows % 128 or cols % 64 or scales.numel() != rows * cols // 16:
raise ValueError("NVFP4 fused dequantization requires exact 128x4 scale geometry")
if dtype not in (torch.bfloat16, torch.float16, torch.float32):
raise ValueError("NVFP4 fused dequantization requires a floating output dtype")
output = torch.empty((rows, cols), dtype=dtype, device=packed.device)
_dequantize_nvfp4[(triton.cdiv(rows * cols, 1024), )](
packed,
scales.view(torch.float8_e4m3fn),
output,
rows,
cols,
# Match Torch's CPU-scalar division: form the
# reciprocal in double, then cast to FP32.
1.0 / global_scale,
BLOCK=1024,
num_warps=4)
return output
|