perf: fp8 rings, lean autocast, gemm staging
- Finalize scale rings inside the quantize kernels: a last-block epilogue (threadfence + counter elect) folds amax into hist, reduces the window and publishes the next scale on device, zero extra launches; _ScaleRing packs [hist | scale | counter] into one CUDA buffer. - Split FP8QuantizeParams out of FP8Params so each operator owns its fields; linear_forward/backward_fp8 take optional ring arguments. - Drop the inference weight-quantization cache; the optimizer bumps the weight version every step, so a cache would miss anyway. - Zero amax scratch via empty + cudaMemsetAsync instead of torch::zeros, cutting a ~50us fill_ dispatch per quantize. - Stage crosswise-B operands K-major with cp.async (contract >= 8192) and PRMT-transpose per k_seg region in smem, interleaved with the MMAs; the sync LDG + byte-scatter path it replaces was long-scoreboard bound (ncu 4.6 vs 0.4 stalls/issue). - Load crosswise-A direct with an in-register PRMT transpose; its operands are typically L2-resident and the staging round trip measured as a net loss. - Enable grouped rasterization for the congruous NT forward (shared B stripe keeps the weight operand hot in L2) and make the smem budget layout-aware (Fp8GemmSmem) while holding two CTAs per SM. - Annotate ops/fp8.py return types; drop weight-cache and decorator tests, hoist their imports to module level. e2e 12L/dim1024/B4xT512 fused AdamW: fp8 137.8ms/step vs bf16 210.3ms, 1.53x. Kernel vs cuBLASLt _scaled_mm: fwd 1.03-1.09x, dX 1.33-1.47x, dW 1.30-1.39x (from 1.10/1.42-1.49/1.52-1.56x), before the pre-transposed copies cuBLASLt needs for dX/dW. fp8 train step vs bf16: 1.34x at 2048 tokens (was 1.25x), 1.08x at 512.
This commit is contained in:
+33
-25
@@ -4,7 +4,7 @@ Isolates the ``fp8_ops`` CUDA extension behind stable Python primitives:
|
||||
|
||||
- ``quantize_bf16(x, scale, fmt) -> (x8, amax)`` — BF16 → FP8 with fused amax
|
||||
- ``mm_fp8(a8, b8, sa, sb) -> out`` — pre-quantized FP8 GEMM (BF16 output)
|
||||
- ``linear_forward_fp8(x, w, bias, sx, sw) -> (out, amax_x, amax_w)``
|
||||
- ``linear_forward_fp8(x, w, bias, sx, sw) -> (out, x8, w8, amax_x, amax_w)``
|
||||
- ``linear_backward_fp8(g, x, w, masks, sg, sw, sx, fmt) -> (gx, gw, gb, amax_g)``
|
||||
|
||||
Scale semantics: scales are *quantization steps* — the value divided out when
|
||||
@@ -16,6 +16,8 @@ Policy (scales, amax history, delayed scaling, autocast) lives in ``fp8.py``;
|
||||
this module is stateless.
|
||||
"""
|
||||
|
||||
from typing import List, Optional, Tuple
|
||||
|
||||
import torch
|
||||
from torch.library import custom_op
|
||||
|
||||
@@ -39,7 +41,7 @@ def _fmt_dtype(fmt: str) -> torch.dtype:
|
||||
@custom_op("custom::fp8_quantize", mutates_args=())
|
||||
def fp8_quantize(
|
||||
x: torch.Tensor, scale: torch.Tensor, fmt: int
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""BF16 -> FP8 quantize with fused amax; returns ``(x8, amax)``."""
|
||||
|
||||
|
||||
@@ -73,7 +75,7 @@ def fp8_gemm(
|
||||
sa: torch.Tensor,
|
||||
sb: torch.Tensor,
|
||||
out_dtype: int = 0,
|
||||
out_scale: torch.Tensor | None = None,
|
||||
out_scale: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
"""FP8 GEMM: ``a @ b * (sa * sb)`` with FP32 accumulation.
|
||||
|
||||
@@ -106,7 +108,9 @@ def _fp8_gemm_cpu(a, b, sa, sb, out_dtype=0, out_scale=None):
|
||||
return acc.to(torch.bfloat16)
|
||||
|
||||
|
||||
def quantize_bf16(x: torch.Tensor, scale: torch.Tensor, fmt: str = "e4m3"):
|
||||
def quantize_bf16(
|
||||
x: torch.Tensor, scale: torch.Tensor, fmt: str = "e4m3"
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""BF16 -> FP8 quantize with fused amax; returns ``(x8, amax)``.
|
||||
|
||||
``scale`` is the quantization step (device scalar); ``fmt`` selects
|
||||
@@ -122,7 +126,7 @@ def mm_fp8(
|
||||
sa: torch.Tensor,
|
||||
sb: torch.Tensor,
|
||||
out_dtype: str = "bf16",
|
||||
out_scale: torch.Tensor | None = None,
|
||||
out_scale: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
"""Pre-quantized FP8 GEMM: ``a @ b * (sa * sb)``.
|
||||
|
||||
@@ -139,24 +143,28 @@ def mm_fp8(
|
||||
|
||||
|
||||
def linear_forward_fp8(
|
||||
x,
|
||||
w,
|
||||
bias,
|
||||
sx,
|
||||
sw,
|
||||
x: torch.Tensor,
|
||||
w: torch.Tensor,
|
||||
bias: Optional[torch.Tensor],
|
||||
sx: torch.Tensor,
|
||||
sw: torch.Tensor,
|
||||
fmt: str = "e4m3",
|
||||
bias_scale=None,
|
||||
x_ring=None,
|
||||
bias_scale: Optional[torch.Tensor] = None,
|
||||
x_ring: Optional[torch.Tensor] = None,
|
||||
x_ring_idx: int = 0,
|
||||
x_ring_margin: int = 0,
|
||||
w_ring=None,
|
||||
w_ring: Optional[torch.Tensor] = None,
|
||||
w_ring_idx: int = 0,
|
||||
w_ring_margin: int = 0,
|
||||
):
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
"""Pure FP8 linear forward: quantize x/w to ``fmt``, pre-quantized GEMM.
|
||||
|
||||
Returns ``(out, amax_x, amax_w)``. ``bias`` may be ``None``. For static
|
||||
fp8 inference, ``w`` and ``bias`` may arrive pre-quantized to ``fmt``
|
||||
Returns ``(out, x8, w8, amax_x, amax_w)`` — the quantized operands are
|
||||
handed back so the policy layer can cache the weight quantization while
|
||||
the weight tensor is unchanged (torch autocast's cached_cast analog).
|
||||
``x8`` is ``[M, K]`` and ``w8`` is ``[N, K]`` (the passed-in ``w`` itself
|
||||
on the pre-quantized path). ``bias`` may be ``None``. For static fp8
|
||||
inference, ``w`` and ``bias`` may arrive pre-quantized to ``fmt``
|
||||
(produced by :func:`quantize_bf16` with their scales as ``sw`` /
|
||||
``bias_scale``); a pre-quantized ``bias`` requires ``bias_scale``, and
|
||||
its ``amax_w`` comes back 0. The bias is fused into the GEMM epilogue.
|
||||
@@ -190,18 +198,18 @@ def linear_forward_fp8(
|
||||
|
||||
|
||||
def linear_backward_fp8(
|
||||
g,
|
||||
x,
|
||||
w,
|
||||
masks,
|
||||
sg,
|
||||
sw,
|
||||
sx,
|
||||
g: torch.Tensor,
|
||||
x: torch.Tensor,
|
||||
w: torch.Tensor,
|
||||
masks: List[bool],
|
||||
sg: torch.Tensor,
|
||||
sw: torch.Tensor,
|
||||
sx: torch.Tensor,
|
||||
fmt: str = "e5m2",
|
||||
g_ring=None,
|
||||
g_ring: Optional[torch.Tensor] = None,
|
||||
g_ring_idx: int = 0,
|
||||
g_ring_margin: int = 0,
|
||||
):
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
"""FP8 linear backward; returns ``(grad_input, grad_weight, grad_bias, amax_g)``.
|
||||
|
||||
The gradient (and the transposed w/x operands) are quantized to ``fmt``
|
||||
|
||||
Reference in New Issue
Block a user