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
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.
fastvideo.layers.quantization.fp8_kernels.scaled_mm_token_channel ¶
(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.