Skip to content

fp8_kernels

Triton helpers for per-token x per-channel FP8 linears on GPUs whose rowwise _scaled_mm is slow.

On sm89 (RTX 4090 / L40S / RTX 6000 Ada) torch._scaled_mm with rowwise scales runs at ~70 TFLOPS, below bf16 (~160), while the per-tensor kernel reaches 220-305 TFLOPS. The per-token x per-channel result is recovered exactly by running the per-tensor kernel with unit scales and applying out[i, j] *= sx[i] * sw[j] in one pass over the output, which costs 5-10% of the GEMM instead of 2-4x.

Functions:

fastvideo.layers.quantization.fp8_kernels.quantize_rowwise_fp8

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

Per-token FP8 quantization in one launch. Returns (x_fp8 [M, K], x_scale [M, 1] float32).

Source code in fastvideo/layers/quantization/fp8_kernels.py
def quantize_rowwise_fp8(x_2d: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    """Per-token FP8 quantization in one launch. Returns ``(x_fp8 [M, K], x_scale [M, 1] float32)``."""
    x_2d = x_2d.contiguous()
    M, K = x_2d.shape
    q = torch.empty((M, K), device=x_2d.device, dtype=torch.float8_e4m3fn)
    s = torch.empty((M, 1), device=x_2d.device, dtype=torch.float32)
    if M:
        _quantize_rowwise_kernel[(M, )](x_2d, q, s, K, x_2d.stride(0), q.stride(0), BLOCK_K=1024, num_warps=4)
    return q, s

fastvideo.layers.quantization.fp8_kernels.rowwise_scaled_mm_is_slow cached

rowwise_scaled_mm_is_slow() -> bool

Ada (sm89) has no fast rowwise-scaled FP8 GEMM in torch; Hopper and Blackwell do.

Source code in fastvideo/layers/quantization/fp8_kernels.py
@functools.cache
def rowwise_scaled_mm_is_slow() -> bool:
    """Ada (sm89) has no fast rowwise-scaled FP8 GEMM in torch; Hopper and Blackwell do."""
    return torch.cuda.is_available() and torch.cuda.get_device_capability() == (8, 9)

fastvideo.layers.quantization.fp8_kernels.scaled_mm_token_channel

scaled_mm_token_channel(x_fp8: Tensor, x_scale: Tensor, w_fp8_t: Tensor, w_scale: Tensor) -> Tensor

(x_fp8 * x_scale) @ (w_fp8_t * w_scale) in bf16 via the fast per-tensor GEMM plus a scale epilogue.

The unit-scale GEMM output is at most 448^2 * K, far inside bf16 range, and its relative precision is that of any bf16 output, so the epilogue loses nothing against rowwise scaling.

Source code in fastvideo/layers/quantization/fp8_kernels.py
def scaled_mm_token_channel(x_fp8: torch.Tensor, x_scale: torch.Tensor, w_fp8_t: torch.Tensor,
                            w_scale: torch.Tensor) -> torch.Tensor:
    """``(x_fp8 * x_scale) @ (w_fp8_t * w_scale)`` in bf16 via the fast per-tensor GEMM plus a scale epilogue.

    The unit-scale GEMM output is at most 448^2 * K, far inside bf16 range, and its relative
    precision is that of any bf16 output, so the epilogue loses nothing against rowwise scaling.
    """
    one = torch.ones((), device=x_fp8.device, dtype=torch.float32)
    out = torch._scaled_mm(x_fp8, w_fp8_t, scale_a=one, scale_b=one, out_dtype=torch.bfloat16)
    if isinstance(out, tuple):
        out = out[0]
    M, N = out.shape
    grid = (triton.cdiv(M, 64), triton.cdiv(N, 128))
    _scale_rows_cols_kernel[grid](out, x_scale.reshape(-1), w_scale.reshape(-1), M, N, BLOCK_M=64, BLOCK_N=128)
    return out