Skip to content

nvfp4_dequant

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