- 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.
242 lines
8.2 KiB
Python
242 lines
8.2 KiB
Python
"""FP8 CUDA kernel interface adapter (the only module touching the pybind).
|
|
|
|
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, 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
|
|
quantizing (``x8 = x / scale``). Every primitive computes its own inverse
|
|
internally; callers never pass ``scale_inv``. ``amax`` values are *returned*,
|
|
never passed as output arguments. ``fmt`` is ``"e4m3"`` or ``"e5m2"``.
|
|
|
|
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
|
|
|
|
from astrai.extension.loader import get_module
|
|
|
|
# fmt string -> kernel int (0 = E4M3, 1 = E5M2)
|
|
_FMT_TO_INT = {"e4m3": 0, "e5m2": 1}
|
|
|
|
|
|
def _fmt_int(fmt: str) -> int:
|
|
try:
|
|
return _FMT_TO_INT[fmt]
|
|
except KeyError:
|
|
raise ValueError(f"unsupported fp8 format {fmt!r} (expected 'e4m3' or 'e5m2')")
|
|
|
|
|
|
def _fmt_dtype(fmt: str) -> torch.dtype:
|
|
return torch.float8_e5m2 if _fmt_int(fmt) else torch.float8_e4m3fn
|
|
|
|
|
|
@custom_op("custom::fp8_quantize", mutates_args=())
|
|
def fp8_quantize(
|
|
x: torch.Tensor, scale: torch.Tensor, fmt: int
|
|
) -> Tuple[torch.Tensor, torch.Tensor]:
|
|
"""BF16 -> FP8 quantize with fused amax; returns ``(x8, amax)``."""
|
|
|
|
|
|
@fp8_quantize.register_fake
|
|
def _fp8_quantize_fake(x, scale, fmt):
|
|
dtype = torch.float8_e5m2 if fmt else torch.float8_e4m3fn
|
|
return (
|
|
torch.empty(x.shape, device=x.device, dtype=dtype),
|
|
torch.empty(1, device=x.device, dtype=torch.float32),
|
|
)
|
|
|
|
|
|
@fp8_quantize.register_kernel("cuda")
|
|
def _fp8_quantize_cuda(x, scale, fmt):
|
|
if x.dtype != torch.bfloat16:
|
|
raise TypeError(f"fp8 quantize requires bf16 input, got {x.dtype}")
|
|
return get_module("fp8_ops").quantize_bf16(x, scale, int(fmt))
|
|
|
|
|
|
@fp8_quantize.register_kernel("cpu")
|
|
def _fp8_quantize_cpu(x, scale, fmt):
|
|
x8 = (x.float() / scale).to(_fmt_dtype("e5m2" if fmt else "e4m3"))
|
|
amax = x.abs().amax().float().reshape(1).clamp_min(1e-12)
|
|
return x8, amax
|
|
|
|
|
|
@custom_op("custom::fp8_gemm", mutates_args=())
|
|
def fp8_gemm(
|
|
a: torch.Tensor,
|
|
b: torch.Tensor,
|
|
sa: torch.Tensor,
|
|
sb: torch.Tensor,
|
|
out_dtype: int = 0,
|
|
out_scale: Optional[torch.Tensor] = None,
|
|
) -> torch.Tensor:
|
|
"""FP8 GEMM: ``a @ b * (sa * sb)`` with FP32 accumulation.
|
|
|
|
``out_dtype``: 0 = BF16 (default), 1 = FP8 E4M3 (requires ``out_scale``,
|
|
the quantization step for the output — mirrors ``torch._scaled_mm``).
|
|
"""
|
|
|
|
|
|
@fp8_gemm.register_fake
|
|
def _fp8_gemm_fake(a, b, sa, sb, out_dtype=0, out_scale=None):
|
|
dtype = torch.float8_e4m3fn if out_dtype else torch.bfloat16
|
|
return torch.empty((a.size(0), b.size(1)), device=a.device, dtype=dtype)
|
|
|
|
|
|
@fp8_gemm.register_kernel("cuda")
|
|
def _fp8_gemm_cuda(a, b, sa, sb, out_dtype=0, out_scale=None):
|
|
if a.dtype != b.dtype or a.dtype not in (torch.float8_e4m3fn, torch.float8_e5m2):
|
|
raise TypeError(
|
|
f"fp8 GEMM requires matching fp8 inputs, got {a.dtype}/{b.dtype}"
|
|
)
|
|
return get_module("fp8_ops").mm_fp8(a, b, sa, sb, int(out_dtype), out_scale)
|
|
|
|
|
|
@fp8_gemm.register_kernel("cpu")
|
|
def _fp8_gemm_cpu(a, b, sa, sb, out_dtype=0, out_scale=None):
|
|
acc = a.float() @ b.float() * sa * sb
|
|
if out_dtype:
|
|
os_ = 1.0 if out_scale is None else out_scale
|
|
return (acc * os_).to(torch.float8_e4m3fn)
|
|
return acc.to(torch.bfloat16)
|
|
|
|
|
|
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
|
|
E4M3 or E5M2. ``amax`` is a fresh 1-element float32 tensor — the caller
|
|
never clears it.
|
|
"""
|
|
return fp8_quantize(x, scale, _fmt_int(fmt))
|
|
|
|
|
|
def mm_fp8(
|
|
a: torch.Tensor,
|
|
b: torch.Tensor,
|
|
sa: torch.Tensor,
|
|
sb: torch.Tensor,
|
|
out_dtype: str = "bf16",
|
|
out_scale: Optional[torch.Tensor] = None,
|
|
) -> torch.Tensor:
|
|
"""Pre-quantized FP8 GEMM: ``a @ b * (sa * sb)``.
|
|
|
|
``a``/``b`` must be FP8 tensors of the same format (E4M3 or E5M2);
|
|
``sa``/``sb`` are their quantization steps. ``out_dtype`` is ``"bf16"``
|
|
(default) or ``"e4m3"`` — FP8 output for layer-to-layer pipelines, which
|
|
requires ``out_scale`` (the output quantization step).
|
|
"""
|
|
if out_dtype not in ("bf16", "e4m3"):
|
|
raise ValueError(
|
|
f"unsupported out_dtype {out_dtype!r} (expected 'bf16' or 'e4m3')"
|
|
)
|
|
return fp8_gemm(a, b, sa, sb, int(out_dtype == "e4m3"), out_scale)
|
|
|
|
|
|
def linear_forward_fp8(
|
|
x: torch.Tensor,
|
|
w: torch.Tensor,
|
|
bias: Optional[torch.Tensor],
|
|
sx: torch.Tensor,
|
|
sw: torch.Tensor,
|
|
fmt: str = "e4m3",
|
|
bias_scale: Optional[torch.Tensor] = None,
|
|
x_ring: Optional[torch.Tensor] = None,
|
|
x_ring_idx: int = 0,
|
|
x_ring_margin: int = 0,
|
|
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, 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.
|
|
``x_ring`` / ``w_ring`` (delayed scaling) are ``[hist | scale | counter]``
|
|
float32 buffers the quantize kernels finalize in-kernel: the measured
|
|
amax lands in ``hist[idx]`` and the next step's scale is published on
|
|
device, replacing the eager hist/max/scale update chain.
|
|
"""
|
|
fmt8 = _fmt_dtype(fmt)
|
|
if x.dtype != torch.bfloat16 or w.dtype not in (torch.bfloat16, fmt8):
|
|
raise TypeError(
|
|
f"fp8 forward requires bf16 x and bf16-or-{fmt} w, got {x.dtype}/{w.dtype}"
|
|
)
|
|
if bias is None:
|
|
bias = torch.empty(0, device=x.device, dtype=x.dtype)
|
|
return get_module("fp8_ops").linear_forward_fp8(
|
|
x,
|
|
w,
|
|
bias,
|
|
sx,
|
|
sw,
|
|
_fmt_int(fmt),
|
|
bias_scale,
|
|
x_ring,
|
|
x_ring_idx,
|
|
x_ring_margin,
|
|
w_ring,
|
|
w_ring_idx,
|
|
w_ring_margin,
|
|
)
|
|
|
|
|
|
def linear_backward_fp8(
|
|
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: 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``
|
|
(default E5M2 — larger dynamic range for gradients) and the two GEMMs run
|
|
as FP8 tensor-core products sharing a single gradient quantization.
|
|
``g_ring`` (delayed scaling) is a ``[hist | scale | counter]`` buffer the
|
|
g quantize kernel finalizes in-kernel (see :func:`linear_forward_fp8`).
|
|
"""
|
|
if not (
|
|
g.dtype == torch.bfloat16
|
|
and x.dtype == torch.bfloat16
|
|
and w.dtype == torch.bfloat16
|
|
):
|
|
raise TypeError(
|
|
f"fp8 backward requires bf16 inputs, got {g.dtype}/{x.dtype}/{w.dtype}"
|
|
)
|
|
return get_module("fp8_ops").linear_backward_fp8(
|
|
g,
|
|
x,
|
|
w,
|
|
list(masks),
|
|
sg,
|
|
sw,
|
|
sx,
|
|
_fmt_int(fmt),
|
|
g_ring,
|
|
g_ring_idx,
|
|
g_ring_margin,
|
|
)
|