refactor: reorganize CUDA kernels into per-family directories

- move attention kernels to csrc/kernels/attention/ and rotary to rotary/
- add shared common/mma.cuh (mma_sync, ldmatrix) and device.cuh (sm checks)
- split fp8_mm into three-layer fp8/common.h, gemm.cuh, mm.cu
- fix fused FP8 GEMM ldmatrix lane indexing to fix OOB shared reads
- update extension ops, loader, and kernel tests
This commit is contained in:
2026-08-22 20:40:31 +08:00
parent cb21af38ba
commit 16a55bb474
30 changed files with 1956 additions and 1235 deletions
+230 -158
View File
@@ -1,161 +1,214 @@
"""FP8 training: scaling state and aten::linear dispatch.
"""FP8 training: scaling recipes, per-tensor state, and aten::linear dispatch.
Layered (see also ``ops/fp8.py`` for the CUDA interface adapter):
1. Kernel interface: ``ops.fp8`` - the only module touching the pybind.
2. Training state (this module): per-tensor scales, amax history, delayed
scaling, and the ``fp8_autocast`` context (TE-style, like
``torch.autocast``).
1. Kernel interface: ``ops.fp8`` the only module touching the pybind.
2. Training state (this module): scaling *recipes* (TE-style delayed scaling
or dynamic current-amax scaling), per-tensor scales + amax history, and
the ``fp8_autocast`` context (like ``torch.autocast``).
3. aten::linear integration (this module): registers the CUDA impl and the
M/N alignment guard.
dtype guard.
Usage::
from astrai.extension.fp8 import fp8_autocast
with fp8_autocast(enabled=True):
with fp8_autocast(enabled=True, fp8_format="hybrid"):
logits = model(input_ids)
loss.backward()
Importing this module registers the aten::linear CUDA implementation.
Format defaults follow the ecosystem consensus: E4M3 for the forward pass,
E5M2 for the backward (gradient) pass ("hybrid"); every operand's scale is a
quantization step derived from its amax history by the active recipe.
"""
from contextlib import contextmanager
from dataclasses import dataclass
from enum import Enum
from typing import Optional
import torch
from torch.library import Library
from astrai.extension.ops.fp8 import (
linear_backward_scaled,
linear_forward_scaled,
linear_backward_fp8,
linear_forward_fp8,
mm_fp8,
quantize_bf16,
)
E4M3_MAX = 448.0
# Max representable value per FP8 format (E4M3: 448, E5M2: 57344).
FP8_MAX = {"e4m3": 448.0, "e5m2": 57344.0}
E4M3_MAX = FP8_MAX["e4m3"] # legacy alias
# ---------------------------------------------------------------------------
# Layer 2: training state (scales, amax history, delayed scaling, autocast)
# ---------------------------------------------------------------------------
class FP8Format(str, Enum):
"""Per-direction FP8 format. HYBRID = E4M3 forward / E5M2 backward."""
E4M3 = "e4m3"
E5M2 = "e5m2"
HYBRID = "hybrid"
def fwd(self) -> str:
return "e4m3" if self is FP8Format.HYBRID else self.value
def bwd(self) -> str:
return "e5m2" if self is FP8Format.HYBRID else self.value
class FP8Recipe:
"""Scale-from-amax policy; the scale computation is the injection point.
``scale_from_history`` receives the amax tensor for this operand (a ring
window for delayed scaling, the current amax for dynamic scaling) and
returns the quantization step: ``scale = (amax / FP8_MAX[fmt]) / 2^margin``.
"""
history_len: int = 16
margin: int = 0
def scale_from_history(self, amax: torch.Tensor, fmt: str) -> torch.Tensor:
raise NotImplementedError
@dataclass
class DelayedScaling(FP8Recipe):
"""TE-style delayed scaling: max over the amax history window.
The scale is computed from amax measured in *previous* steps (delayed one
step); the window length trades responsiveness against stability.
"""
history_len: int = 16
margin: int = 0
def scale_from_history(self, amax: torch.Tensor, fmt: str) -> torch.Tensor:
peak = amax.max()
return ((peak / FP8_MAX[fmt]) / (2**self.margin)).clamp_min(1e-12)
@dataclass
class DynamicScaling(FP8Recipe):
"""Current-amax scaling (torchao DYNAMIC): measure, then quantize.
No history — the scale is derived from the amax of the tensor being
quantized in the same step, at the cost of an extra reduction pass.
"""
history_len: int = 1
margin: int = 0
def scale_from_history(self, amax: torch.Tensor, fmt: str) -> torch.Tensor:
peak = amax.max()
return ((peak / FP8_MAX[fmt]) / (2**self.margin)).clamp_min(1e-12)
class FP8TensorMeta:
"""Scales + amax state for one weight tensor and its paired activations.
"""Per-tensor scaling state: amax history rings + derived scales.
- weight: delayed scale from a 16-step amax history window (TE style)
- x/g: delayed one step, reuse the quantize kernel's free atomic amax
One ring per operand (weight / activation / gradient). Scales are derived
from the ring by the recipe; fused kernels record the amax while
quantizing, so the scale used at step N reflects amax from steps < N
(delayed one step).
"""
__slots__ = (
"scale",
"scale_inv",
"amax_history",
"idx",
"x_scale",
"x_scale_inv",
"x_history",
"recipe",
"w_hist",
"x_hist",
"g_hist",
"w_idx",
"x_idx",
"g_scale",
"g_scale_inv",
"g_history",
"g_idx",
"w_scale",
"x_scale",
"g_scale",
"w_init",
"x_init",
"g_init",
)
def __init__(self, device: torch.device, update_interval: int):
self.scale = torch.ones(1, device=device, dtype=torch.float32)
self.scale_inv = torch.ones(1, device=device, dtype=torch.float32)
self.amax_history = torch.ones(
update_interval, device=device, dtype=torch.float32
)
self.idx = 0
def __init__(self, device: torch.device, recipe: FP8Recipe):
self.recipe = recipe
n = recipe.history_len
self.w_hist = torch.ones(n, device=device, dtype=torch.float32)
self.x_hist = torch.ones(n, device=device, dtype=torch.float32)
self.g_hist = torch.ones(n, device=device, dtype=torch.float32)
self.w_idx = self.x_idx = self.g_idx = 0
self.w_scale = torch.ones(1, device=device, dtype=torch.float32)
self.x_scale = torch.ones(1, device=device, dtype=torch.float32)
self.x_scale_inv = torch.ones(1, device=device, dtype=torch.float32)
self.x_history = torch.ones(update_interval, device=device, dtype=torch.float32)
self.x_idx = 0
self.g_scale = torch.ones(1, device=device, dtype=torch.float32)
self.g_scale_inv = torch.ones(1, device=device, dtype=torch.float32)
self.g_history = torch.ones(update_interval, device=device, dtype=torch.float32)
self.g_idx = 0
self.w_init = False
self.x_init = False
self.g_init = False
self.w_init = self.x_init = self.g_init = False
def init_scale(self, t: torch.Tensor) -> None:
"""Immediate scale from the current amax; used on the first call.
# -- ring helpers -------------------------------------------------------
A scale of 1 would underflow small activations/gradients (e4m3 min
normal is 2^-6); initialize from the actual amax once, then delayed
updates take over.
"""
def _record(self, hist: torch.Tensor, idx: int, amax: torch.Tensor) -> int:
hist[idx] = amax.reshape(())
return (idx + 1) % hist.numel()
def _refresh(self, hist: torch.Tensor, scale: torch.Tensor, fmt: str) -> None:
scale.copy_(self.recipe.scale_from_history(hist, fmt))
def _seed(
self, hist: torch.Tensor, scale: torch.Tensor, t: torch.Tensor, fmt: str
) -> None:
amax = t.abs().amax().to(torch.float32).clamp_min(1e-12)
self.scale.copy_(amax / E4M3_MAX)
self.scale_inv.copy_(E4M3_MAX / amax)
self.record(amax)
hist.fill_(amax)
scale.copy_(self.recipe.scale_from_history(hist, fmt))
def push_x_scale(self, amax: torch.Tensor) -> None:
"""Window update for the activation scale (delayed, TE style)."""
self.x_history[self.x_idx] = amax.reshape(())
self.x_idx = (self.x_idx + 1) % self.x_history.numel()
m = self.x_history.max()
self.x_scale.copy_(m / E4M3_MAX)
self.x_scale_inv.copy_(E4M3_MAX / m)
# -- per-operand updates (delayed: record now, refresh for next step) ---
def push_g_scale(self, amax: torch.Tensor) -> None:
"""Window update for the gradient scale (delayed, TE style)."""
self.g_history[self.g_idx] = amax.reshape(())
self.g_idx = (self.g_idx + 1) % self.g_history.numel()
m = self.g_history.max()
self.g_scale.copy_(m / E4M3_MAX)
self.g_scale_inv.copy_(E4M3_MAX / m)
def update_w(self, amax: torch.Tensor, fmt: str) -> None:
self.w_idx = self._record(self.w_hist, self.w_idx, amax)
self._refresh(self.w_hist, self.w_scale, fmt)
def record(self, amax: torch.Tensor) -> None:
"""Push the latest amax into the ring buffer (device-side copy, no sync)."""
self.amax_history[self.idx] = amax.reshape(())
self.idx = (self.idx + 1) % self.amax_history.numel()
def update_x(self, amax: torch.Tensor, fmt: str) -> None:
self.x_idx = self._record(self.x_hist, self.x_idx, amax)
self._refresh(self.x_hist, self.x_scale, fmt)
def refresh(self) -> None:
"""Recompute scale from the amax history window (delayed scaling)."""
amax = self.amax_history.max()
if amax > 0:
self.scale.copy_(amax / E4M3_MAX)
self.scale_inv.copy_(E4M3_MAX / amax)
def update_g(self, amax: torch.Tensor, fmt: str) -> None:
self.g_idx = self._record(self.g_hist, self.g_idx, amax)
self._refresh(self.g_hist, self.g_scale, fmt)
# -- first-use seeding --------------------------------------------------
def init_w(self, w: torch.Tensor, fmt: str) -> None:
self._seed(self.w_hist, self.w_scale, w, fmt)
self.w_init = True
def init_x(self, x: torch.Tensor, fmt: str) -> None:
self._seed(self.x_hist, self.x_scale, x, fmt)
self.x_init = True
def init_g(self, g: torch.Tensor, fmt: str) -> None:
self._seed(self.g_hist, self.g_scale, g, fmt)
self.g_init = True
class FP8State:
"""Global fp8 training state, TE-style."""
"""Global fp8 training state: active recipe + per-tensor metas."""
def __init__(self, update_interval: int = 16):
def __init__(self):
self.enabled = False
self.update_interval = update_interval
self.step_count = 0
self.recipe: FP8Recipe = DelayedScaling()
self.fp8_format: FP8Format = FP8Format.HYBRID
self._metas: dict[tuple, FP8TensorMeta] = {}
self._last_device: torch.device | None = None
def _get_device(self, t: torch.Tensor) -> torch.device:
if self._last_device is None:
self._last_device = t.device
return t.device
self._last_device: Optional[torch.device] = None
def get_weight_meta(self, w: torch.Tensor) -> FP8TensorMeta:
key = (w.data_ptr(), w.shape, w.dtype)
meta = self._metas.get(key)
if meta is None:
meta = FP8TensorMeta(self._get_device(w), self.update_interval)
if self._last_device is None:
self._last_device = w.device
meta = FP8TensorMeta(w.device, self.recipe)
self._metas[key] = meta
return meta
def step(self) -> None:
"""Advance the counter and refresh all weight scales every N steps."""
self.step_count += 1
if self.step_count % self.update_interval == 0:
for meta in self._metas.values():
meta.refresh()
def reset(self) -> None:
self.enabled = False
self.step_count = 0
self._metas.clear()
self._last_device = None
@@ -171,100 +224,114 @@ def fp8_state() -> FP8State:
@contextmanager
def fp8_autocast(enabled: bool = True, update_interval: int = 16):
def fp8_autocast(
enabled: bool = True,
update_interval: int = 16,
recipe: Optional[FP8Recipe] = None,
fp8_format: str = "hybrid",
margin: int = 0,
):
"""Autocast-style context: fp8 linear dispatch on this thread.
Usage::
with fp8_autocast(enabled=True):
with fp8_autocast(enabled=True, fp8_format="hybrid"):
logits = model(input_ids) # aten::linear -> fp8 path
loss.backward()
The scale-update counter advances once per ``enter`` (one training step),
refreshing weight scales from their amax history every ``update_interval``.
Args:
enabled: toggle fp8 dispatch for aten::linear.
update_interval: legacy alias for the delayed-scaling history window
(used only when ``recipe`` is not given).
recipe: scaling policy; defaults to ``DelayedScaling(update_interval)``.
fp8_format: ``"e4m3"`` / ``"e5m2"`` / ``"hybrid"`` (default) — hybrid
means E4M3 forward, E5M2 backward.
margin: scale headroom (``scale = (amax / FP8_MAX) / 2^margin``) used
with the default delayed recipe.
"""
state = fp8_state()
prev_enabled = state.enabled
prev_interval = state.update_interval
prev = (state.enabled, state.recipe, state.fp8_format)
if recipe is None:
recipe = DelayedScaling(history_len=update_interval, margin=margin)
state.enabled = enabled
state.update_interval = update_interval
state.recipe = recipe
state.fp8_format = FP8Format(fp8_format)
try:
if enabled:
state.step()
yield
finally:
state.enabled = prev_enabled
state.update_interval = prev_interval
state.enabled, state.recipe, state.fp8_format = prev
# ---------------------------------------------------------------------------
# Strategy-level forward / backward (called from the aten::linear impl)
# ---------------------------------------------------------------------------
def _dynamic_scale(t: torch.Tensor, recipe: FP8Recipe, fmt: str) -> torch.Tensor:
amax = t.abs().amax().to(torch.float32).clamp_min(1e-12)
return recipe.scale_from_history(amax, fmt)
def fp8_linear_forward(x: torch.Tensor, w: torch.Tensor, bias=None):
"""TE-style scaled fp8 linear forward (called from the aten::linear impl).
"""Scaled fp8 linear forward (called from the aten::linear impl).
x uses the delayed scale of its paired weight meta (amax from the previous
forward of this linear); the quantize kernel emits the current amax for the
next step. No extra abs/max reduce.
Delayed scaling uses the fused BF16->E4M3 GEMM (quantize + amax inside the
kernel); dynamic scaling measures the current amax first and runs the
pre-quantized path.
"""
if bias is None:
bias = torch.empty(0, device=x.device, dtype=x.dtype)
state = fp8_state()
fmt = state.fp8_format.fwd()
meta = state.get_weight_meta(w)
if not meta.w_init:
meta.init_scale(w)
meta.w_init = True
meta.init_w(w, fmt)
if isinstance(state.recipe, DynamicScaling):
x_2d = x.reshape(-1, w.size(1))
sx = _dynamic_scale(x_2d, state.recipe, fmt)
sw = _dynamic_scale(w, state.recipe, fmt)
x8, _ = quantize_bf16(x_2d, sx, fmt)
w8, _ = quantize_bf16(w, sw, fmt)
out = mm_fp8(x8, w8, sx, sw)
out = out.reshape(*x.shape[:-1], w.size(0))
if bias.numel():
out = out + bias
return out
if not meta.x_init:
amax = x.abs().amax().to(torch.float32).clamp_min(1e-12)
meta.x_history.fill_(amax)
meta.x_scale.copy_(amax / E4M3_MAX)
meta.x_scale_inv.copy_(E4M3_MAX / amax)
meta.x_init = True
amax_x = torch.empty(1, device=x.device, dtype=torch.float32)
amax_w = torch.empty(1, device=x.device, dtype=torch.float32)
out = linear_forward_scaled(
x,
w,
bias,
meta.x_scale,
meta.scale,
meta.x_scale_inv,
meta.scale_inv,
amax_x,
amax_w,
)
meta.record(amax_w)
meta.push_x_scale(amax_x)
meta.init_x(x, fmt)
out, amax_x, amax_w = linear_forward_fp8(x, w, bias, meta.x_scale, meta.w_scale)
meta.update_x(amax_x, fmt)
meta.update_w(amax_w, fmt)
return out
def fp8_linear_backward(g, x, w, masks):
"""TE-style scaled fp8 linear backward (called from aten::linear_backward)."""
def fp8_linear_backward(g: torch.Tensor, x: torch.Tensor, w: torch.Tensor, masks):
"""Scaled fp8 linear backward (called from aten::linear_backward).
The gradient is quantized to the backward format (E5M2 in hybrid mode)
and the dX / dW GEMMs share that single quantization.
"""
state = fp8_state()
fmt = state.fp8_format.bwd()
meta = state.get_weight_meta(w)
if not meta.g_init:
amax = g.abs().amax().to(torch.float32).clamp_min(1e-12)
meta.g_history.fill_(amax)
meta.g_scale.copy_(amax / E4M3_MAX)
meta.g_scale_inv.copy_(E4M3_MAX / amax)
meta.g_init = True
amax_g = torch.empty(1, device=g.device, dtype=torch.float32)
out = linear_backward_scaled(
g,
x,
w,
masks,
meta.g_scale,
meta.scale,
meta.x_scale,
meta.g_scale_inv,
meta.scale_inv,
meta.x_scale_inv,
amax_g,
if isinstance(state.recipe, DynamicScaling):
sg = _dynamic_scale(g, state.recipe, fmt)
sw = _dynamic_scale(w, state.recipe, fmt)
sx = _dynamic_scale(x, state.recipe, fmt)
else:
if not meta.g_init:
meta.init_g(g, fmt)
sg, sw, sx = meta.g_scale, meta.w_scale, meta.x_scale
grad_x, grad_w, grad_b, amax_g = linear_backward_fp8(
g, x, w, masks, sg, sw, sx, fmt
)
meta.push_g_scale(amax_g)
return out
if not isinstance(state.recipe, DynamicScaling):
meta.update_g(amax_g, fmt)
return grad_x, grad_w, grad_b
# ---------------------------------------------------------------------------
# Layer 3: aten::linear integration
# aten::linear integration
# ---------------------------------------------------------------------------
@@ -279,9 +346,14 @@ def fp8_linear_enabled() -> bool:
def _fp8_supported(x: torch.Tensor, w: torch.Tensor) -> bool:
"""cuBLASLt fp8 requires M % 16 == 0 and N % 16 == 0 (K is padded)."""
m = x.numel() // x.size(-1)
return m % 16 == 0 and w.size(0) % 16 == 0
"""Shape guard for the fp8 linear path.
Unlike a strict 16-alignment requirement, the fp8 kernels handle unaligned
M/N via boundary checks (slower but correct) — so no whole-call bf16
fallback for small decode batches. Only the K-dimension contraction must
match, and the weight must be 2D.
"""
return x.dim() >= 2 and w.dim() == 2 and x.size(-1) == w.size(1)
def _linear_cuda_impl(x: torch.Tensor, w: torch.Tensor, bias=None):