Skip to content

fp8_quant_kernels

Fused FP8 (e4m3fn) activation quantization for the generic FP8 linear path.

The eager formulas in fp8_config (abs -> amax -> div -> clamp -> to(float8)) are four to five memory passes over the activation. At MiniMax-H3's shapes that costs about as much as the FP8 _scaled_mm it feeds (1.06 ms next to a 1.25 ms GEMM for a 65,536 x 5,376 activation on MI355X). These kernels do the same math in one pass (rowwise) or an aminmax reduction plus one pass (tensorwise): 3x (tensorwise) and 4-5x (rowwise) faster at those shapes.

NaN handling follows the eager formulas: a NaN anywhere in a row (rowwise) or in the tensor (tensorwise) makes its scale and its FP8 values NaN. Triton's default maximum/minimum and tl.max return the non-NaN operand, so every max, min and clamp below asks for NaN propagation.

Functions:

fastvideo.layers.quantization.fp8_quant_kernels.quantize_rowwise_fused

quantize_rowwise_fused(x_2d: Tensor) -> tuple[Tensor, Tensor]

Returns (x_fp8 [M, K], x_scale [M, 1] float32); one kernel launch.

Source code in fastvideo/layers/quantization/fp8_quant_kernels.py
def quantize_rowwise_fused(x_2d: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    """Returns ``(x_fp8 [M, K], x_scale [M, 1] float32)``; one kernel launch."""
    x_2d = x_2d.contiguous()
    m, k = x_2d.shape
    out = torch.empty((m, k), dtype=FP8_DTYPE, device=x_2d.device)
    scale = torch.empty((m, 1), dtype=torch.float32, device=x_2d.device)
    block = min(4096, triton.next_power_of_2(k))
    if m:
        _rowwise_quant_kernel[(m, )](x_2d,
                                     out,
                                     scale,
                                     k,
                                     x_2d.stride(0),
                                     out.stride(0),
                                     FP8_MAX=FP8_MAX,
                                     MIN_SCALE=FP8_MIN_SCALE,
                                     BLOCK=block,
                                     num_warps=8)
    return out, scale

fastvideo.layers.quantization.fp8_quant_kernels.quantize_tensorwise_fused

quantize_tensorwise_fused(x_2d: Tensor) -> tuple[Tensor, Tensor]

Returns (x_fp8 [M, K], x_scale [1] float32); an aminmax reduction + one pass.

Source code in fastvideo/layers/quantization/fp8_quant_kernels.py
def quantize_tensorwise_fused(x_2d: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    """Returns ``(x_fp8 [M, K], x_scale [1] float32)``; an ``aminmax`` reduction + one pass."""
    x_2d = x_2d.contiguous()
    # aminmax, maximum and clamp all propagate NaN, as the eager amax does.
    lo, hi = torch.aminmax(x_2d)
    amax = torch.maximum(lo.abs(), hi.abs()).float()
    scale = (amax / FP8_MAX).clamp(min=FP8_MIN_SCALE).view(1)
    inv = 1.0 / scale
    out = torch.empty_like(x_2d, dtype=FP8_DTYPE)
    n = x_2d.numel()
    block = 4096
    _scale_cast_kernel[(triton.cdiv(n, block), )](x_2d, out, inv, n, FP8_MAX=FP8_MAX, BLOCK=block, num_warps=8)
    return out, scale