7 Commits
Author SHA1 Message Date
ViperEkura 432dfec3c2 refactor: collapse fp8 recipe hierarchy and state property layers
- Merge DelayedScaling/DynamicScaling and the abstract FP8Recipe base into one FP8Recipe dataclass with a dynamic flag; dispatch now reads cfg.recipe.dynamic instead of isinstance checks
- Drop the _ActiveOrDefault descriptor and the FP8State property views; the persistent defaults are plain default_* attributes and get_weight_meta takes the active recipe explicitly
- Convert FP8TensorMeta to a NamedTuple of the three per-operand rings
- Update tests to the new API; the autocast context test now asserts _active_config push/restore directly
2026-08-31 14:24:51 +08:00
ViperEkura e3c3e28a11 docs: fix stale developer documentation claims
- Move task_alloc/task_free/task_extend/task_cached/task_record_hashes and bind from the PagePool card to a new TaskCacheManager card matching pool.py
- Drop the nonexistent Executor tokenizer attribute and association, add task_cache instead
- Add AllocationStrategy/ContiguousStrategy/PagedStrategy cards and point Allocator/RadixCache composition at PagedStrategy
- Add TaskCacheManager and the allocation strategies to the module overview, add _task_cache to InferenceScheduler
- Fix the design-pattern count in the table of contents (15 -> 16)
- Rewrite the FlashAttnBackend class docstring: packed decode gathers flat K/V via req_to_token and calls flash_attn_varlen_func; dense prefill uses flash_attn_func (no flash_attn_with_kvcache exists)
- Apply the same correction to the backend bullets in internals.md and cuda_kernels.md
- Rename the stale fp8_mma_test.cu reference to fp8_test.cu in cuda_kernels.md
2026-08-31 14:24:51 +08:00
ViperEkura a7d4cb25c5 docs: scope trainer environment variables per job
- Add a Per-Job Environment section explaining that runtime.environment reaches only the GPUs declared in the same job YAML, with one-YAML-per-GPU-group examples for local, cross-PCIe workaround, and NVSwitch NVLink tuning setups
- Replace the NCCL workaround pair in the runtime schema example with ASTR_LOG_LEVEL and ASTR_BACKEND and document value semantics (str() rendering, null exports empty, no host-shell passthrough)
- Comment out the blanket NCCL exports in the get-started multi-GPU example so they are opt-in per docs/guides/distributed.md
- Add a hard rule against copying NCCL workarounds into every training config
2026-08-31 14:24:51 +08:00
ViperEkura 0546331637 fix: skip gradient checkpointing log when no modules configured
- GradientCheckpointingCallback.on_train_begin returns early on empty module list
- previously logged "Gradient checkpointing enabled" even when checkpointing was inactive, misleading profiling
2026-08-31 14:24:51 +08:00
ViperEkura 962c10c52b perf: fold the delayed-scaling ring update into the quantize kernel
- the kernel's last block folds amax into the history window and publishes the next scale in-kernel (atomicAdd ticket + fences), replacing the host update chain
- quantize bindings split into quantize(transposed) / quantize_dual with fixed arities and a QuantLayout enum; the python adapter becomes a thin attention-style wrapper over pybind (Optional ring_state at the boundary, no torch.library custom_ops)
- tests: in-kernel fold vs host reference (exact), dual/transposed orientation byte-equality

Benchmark: L20 (sm_89), 1.2B model, full train step. Per-linear fixed overhead 28.8us -> 8.8us; fp8 vs bf16: M=512 77.5ms, M=2048 144.5ms (1.15x), M=8192 527.4ms (1.28x); losses bit-identical.
2026-08-31 14:24:51 +08:00
ViperEkura 1cf7d6c76b perf: fill steady-state decode input ids via d2d copy
- add InferenceWorkspace.fill_input_ids_from_device copying device tokens straight into the fixed-address input_ids buffer
- cache each decode step's sampled tokens on-device in DecodeSteadyState.last_tokens; when the task signature is unchanged the next step reuses them, replacing the tolist -> python list -> elementwise host fill -> pageable h2d round-trip
- _sample_logits returns (host payload, device tokens); prefill discards the device tensor
- signature change (task join/leave/first decode) still takes the host path; both dispatch paths covered by tests

Benchmark: NVIDIA L20, BF16, 1B model + 0.11B test model (4 layers, hidden 512), contiguous KV cache, CUDA Graph, greedy, prompt 512, generation 256, engine decode via scripts/tools/benchmark.py (alternating A/B, 2-4 paired runs)
- 0.11B batch 32: 21429 -> 24415 tok/s mean (1.14x, +13.9%), 4/4 paired runs faster
- 1B batch 32: 4242 -> 4388 tok/s (1.034x, +3.4%), 7.54 -> 7.29 ms/step
- batch 1: no measurable change (<0.5%)
2026-08-31 14:24:51 +08:00
ViperEkura 36e39496d4 perf: vectorize tiled fp8 transpose quantize and arm amax via memset
- tiled transpose quantize becomes one 64x32-tile kernel: native pair loads (128B warp reads) with in-kernel scalar fallback at unaligned or ragged rows, so odd widths and misaligned bases no longer route to a separate kernel
- the old 32x32 scalar tiled kernel and its launcher correctness branch are gone; grid sizing simplifies to 1 + total / (vec * threads) since both elementwise loops are grid-stride
- quantize arms the amax buffer with cudaMemsetAsync instead of the zeros() fill kernel, dropping one tensor-op dispatch and kernel launch per call
- byte-exact parity holds over 1404 golden records (13 shapes x 3 dtypes x 3 scales x 2 formats x 3 layouts x aligned/misaligned) and tests/extension passes 65/65
- elementwise quantize kernel left unchanged: 16B-store pairing, __ldcs streaming hints and amax tree reduction all measured neutral at its ~52% DRAM ceiling and were reverted

Benchmark: L20 (sm_89), profiler kernel time with L2 flushed between calls.
- transposed quantize (layout 1): 230 -> 294 GB/s on 2048x1536 (+28%), 245 -> 299 on 2048x1536 weights (+22%); dual-layout (layout 2) 248 -> 329 (+33%) on the same shapes
- DRAM-saturated sizes (~10.6M elements) regress ~5% (404 -> 384 GB/s on 8192x1536), ~0.02% of a training step; accepted for the single-kernel shape after scalar-path and geometry variants both measured the same
- amax init fill kernel 3.0us -> memset 0.9us; quantize call CPU wall 18.5 -> 13.4us on 128x1536
2026-08-31 14:24:51 +08:00
16 changed files with 712 additions and 492 deletions
+6 -5
View File
@@ -694,12 +694,13 @@ class CudaBackend(AttentionBackend):
class FlashAttnBackend(AttentionBackend): class FlashAttnBackend(AttentionBackend):
"""FlashAttention backend via the optional ``flash-attn`` package. """FlashAttention backend via the optional ``flash-attn`` package.
Decode (q_len=1, contiguous cache): uses ``flash_attn_with_kvcache``, Decode (q_len=1, contiguous cache): writes K/V to the pool, gathers
which reads K/V directly from the flat pool via cache_batch_idx + flat K/V via the ``req_to_token`` page table, and calls
cache_seqlens — no materialized KV gather. ``flash_attn_varlen_func`` over the ragged batch
(``qo_indptr``/``kv_indptr``).
Prefill / non-contiguous decode: falls back to KV gather + Prefill: packed 3-D calls share the ``flash_attn_varlen_func`` path;
``flash_attn_func``. dense 4-D calls go through ``flash_attn_func`` (mask-free only).
""" """
@classmethod @classmethod
+71 -102
View File
@@ -31,12 +31,12 @@ import functools
from contextvars import ContextVar, Token from contextvars import ContextVar, Token
from dataclasses import dataclass from dataclasses import dataclass
from enum import Enum from enum import Enum
from typing import Dict, List, Optional from typing import Dict, List, NamedTuple, Optional
import torch import torch
from torch.library import Library from torch.library import Library
from astrai.extension.ops.fp8 import mm_fp8, quantize from astrai.extension.ops.fp8 import mm_fp8, quantize, quantize_dual
# Max representable value per FP8 format (E4M3: 448, E5M2: 57344). # Max representable value per FP8 format (E4M3: 448, E5M2: 57344).
FP8_MAX = {"e4m3": 448.0, "e5m2": 57344.0} FP8_MAX = {"e4m3": 448.0, "e5m2": 57344.0}
@@ -56,46 +56,35 @@ class FP8Format(str, Enum):
return "e5m2" if self is FP8Format.HYBRID else self.value return "e5m2" if self is FP8Format.HYBRID else self.value
@dataclass
class FP8Recipe: class FP8Recipe:
"""Scale-from-amax policy: ``scale = (amax / FP8_MAX[fmt]) / 2^margin``. """Scale-from-amax policy: ``scale = (amax / FP8_MAX[fmt]) / 2^margin``.
``scale_from_history`` receives the operand's amax tensor (a ring window for ``dynamic=False`` (default) is TE-style delayed scaling: max over the
delayed scaling, the current amax for dynamic scaling) and returns the amax history window (amax from *previous* steps; the window trades
quantization step. Subclasses set ``history_len`` / ``margin``. responsiveness against stability). ``dynamic=True`` is current-amax
scaling (torchao DYNAMIC): measure, then quantize — no history, at an
extra pass. ``scale_from_history`` receives the operand's amax tensor
(a ring window / the current amax) and returns the quantization step.
""" """
history_len: int = 16 history_len: int = 16
margin: int = 0 margin: int = 0
dynamic: bool = False
def scale_from_history(self, amax: torch.Tensor, fmt: str) -> torch.Tensor: def scale_from_history(self, amax: torch.Tensor, fmt: str) -> torch.Tensor:
peak = amax.max() peak = amax.max()
return ((peak / FP8_MAX[fmt]) / (2**self.margin)).clamp_min(1e-12) return ((peak / FP8_MAX[fmt]) / (2**self.margin)).clamp_min(1e-12)
@dataclass
class DelayedScaling(FP8Recipe):
"""TE-style delayed scaling: max over the amax history window (amax from
*previous* steps; the window trades responsiveness against stability)."""
history_len: int = 16
margin: int = 0
@dataclass
class DynamicScaling(FP8Recipe):
"""Current-amax scaling (torchao DYNAMIC): measure, then quantize. No
history — the scale is derived from the same-step amax, at an extra pass."""
history_len: int = 1
margin: int = 0
class _ScaleRing: class _ScaleRing:
"""One operand's delayed-scaling state: a float32 buffer """One operand's delayed-scaling state: a float32 buffer
``[hist[n] | scale | counter]`` (views). ``update`` folds the amax ``[hist[n] | scale | legacy | amax | done]`` (views). The quantize
returned by the quantize primitive into ``hist[idx]`` and publishes the kernel folds its fused amax into ``hist[idx]`` and publishes the next
next scale from the window; ``idx`` advances host-side each step. The scale from the window in its own last block (``fold_args`` passes the
trailing slot is a legacy counter kept for state-buffer compatibility. buffer + recipe constants); ``idx`` advances host-side each use. The
``amax``/``done`` tail slots are kernel scratch (self-cleaning across
launches); the legacy slot keeps state-buffer compatibility.
""" """
__slots__ = ("recipe", "state", "hist", "scale", "idx", "initialized") __slots__ = ("recipe", "state", "hist", "scale", "idx", "initialized")
@@ -103,7 +92,7 @@ class _ScaleRing:
def __init__(self, device: torch.device, recipe: FP8Recipe): def __init__(self, device: torch.device, recipe: FP8Recipe):
self.recipe = recipe self.recipe = recipe
n = recipe.history_len n = recipe.history_len
self.state = torch.zeros(n + 2, device=device, dtype=torch.float32) self.state = torch.zeros(n + 4, device=device, dtype=torch.float32)
self.hist = self.state[:n] self.hist = self.state[:n]
self.scale = self.state[n : n + 1] self.scale = self.state[n : n + 1]
self.idx = 0 self.idx = 0
@@ -119,23 +108,25 @@ class _ScaleRing:
self.scale.copy_(self.recipe.scale_from_history(self.hist, fmt)) self.scale.copy_(self.recipe.scale_from_history(self.hist, fmt))
self.initialized = True self.initialized = True
def update(self, amax: torch.Tensor, fmt: str) -> None: def fold_args(self, fmt: str) -> dict:
self.hist[self.idx].copy_(amax.reshape(())) """Keyword arguments for quantize()'s in-kernel history fold."""
self.scale.copy_(self.recipe.scale_from_history(self.hist, fmt)) return {
"ring_state": self.state,
"hist_idx": self.idx,
"fp8_max": FP8_MAX[fmt],
"pow2_margin": float(2**self.recipe.margin),
}
class FP8TensorMeta: class FP8TensorMeta(NamedTuple):
"""Per-weight delayed-scaling state for ``w``, ``x`` and ``g``. """Per-weight delayed-scaling rings for ``w``, ``x`` and ``g``.
DynamicScaling never allocates a meta; it measures the current amax inline. Dynamic scaling never allocates a meta; it measures the current amax inline.
""" """
__slots__ = ("w", "x", "g") w: _ScaleRing
x: _ScaleRing
def __init__(self, device: torch.device, recipe: FP8Recipe): g: _ScaleRing
self.w = _ScaleRing(device, recipe)
self.x = _ScaleRing(device, recipe)
self.g = _ScaleRing(device, recipe)
@dataclass(frozen=True) @dataclass(frozen=True)
@@ -159,55 +150,29 @@ _active_config: ContextVar[Optional[_ActiveConfig]] = ContextVar(
class FP8State: class FP8State:
"""Global fp8 training state: per-tensor metas + out-of-region defaults. """Global fp8 training state: per-tensor metas + out-of-region defaults.
The active ``(enabled, recipe, fp8_format)`` triple is a ``ContextVar`` set The active ``(enabled, recipe, fp8_format)`` triple is a ``ContextVar``
by ``fp8_autocast``. The properties below read that active config when a set by ``fp8_autocast`` (see ``_active``/``_current_config``); these plain
region is open and the global defaults otherwise; the setters (and attributes are the persistent defaults applied outside any region —
``fp8_linear_enable``) write the global defaults — the persistent switch ``fp8_linear_enable`` writes ``default_enabled``. The metas registry is
applying outside any region. The metas registry is shared across threads shared across threads (GIL-protected); fp8 backward runs on autograd
(GIL-protected); fp8 backward runs on autograd engine threads and only engine threads and only touches metas captured on ``ctx`` at forward time.
touches metas captured on ``ctx`` at forward time.
""" """
def __init__(self): def __init__(self):
self.default_enabled = False self.default_enabled = False
self.default_recipe: FP8Recipe = DelayedScaling() self.default_recipe: FP8Recipe = FP8Recipe()
self.default_format: FP8Format = FP8Format.HYBRID self.default_format: FP8Format = FP8Format.HYBRID
self._metas: Dict[tuple, FP8TensorMeta] = {} self._metas: Dict[tuple, FP8TensorMeta] = {}
# Active-config views (region config if open, else the defaults). def get_weight_meta(self, w: torch.Tensor, recipe: FP8Recipe) -> FP8TensorMeta:
@property
def enabled(self) -> bool:
cfg = _active_config.get()
return cfg.enabled if cfg is not None else self.default_enabled
@property
def recipe(self) -> FP8Recipe:
cfg = _active_config.get()
return cfg.recipe if cfg is not None else self.default_recipe
@property
def fp8_format(self) -> FP8Format:
cfg = _active_config.get()
return cfg.fp8_format if cfg is not None else self.default_format
# Persistent (out-of-region) defaults.
@enabled.setter
def enabled(self, value: bool) -> None:
self.default_enabled = bool(value)
@recipe.setter
def recipe(self, value: FP8Recipe) -> None:
self.default_recipe = value
@fp8_format.setter
def fp8_format(self, value: FP8Format) -> None:
self.default_format = FP8Format(value)
def get_weight_meta(self, w: torch.Tensor) -> FP8TensorMeta:
key = (w.data_ptr(), w.shape, w.dtype) key = (w.data_ptr(), w.shape, w.dtype)
meta = self._metas.get(key) meta = self._metas.get(key)
if meta is None: if meta is None:
meta = FP8TensorMeta(w.device, self.recipe) meta = FP8TensorMeta(
_ScaleRing(w.device, recipe),
_ScaleRing(w.device, recipe),
_ScaleRing(w.device, recipe),
)
self._metas[key] = meta self._metas[key] = meta
return meta return meta
@@ -215,7 +180,7 @@ class FP8State:
"""Restore construction defaults (switch, recipe, format) and drop all """Restore construction defaults (switch, recipe, format) and drop all
per-weight metas — a full state reset for tests / reconfiguration.""" per-weight metas — a full state reset for tests / reconfiguration."""
self.default_enabled = False self.default_enabled = False
self.default_recipe = DelayedScaling() self.default_recipe = FP8Recipe()
self.default_format = FP8Format.HYBRID self.default_format = FP8Format.HYBRID
self._metas.clear() self._metas.clear()
@@ -278,7 +243,7 @@ class fp8_autocast:
margin: int = 0, margin: int = 0,
): ):
if recipe is None: if recipe is None:
recipe = DelayedScaling(history_len=update_interval, margin=margin) recipe = FP8Recipe(history_len=update_interval, margin=margin)
self._config = _ActiveConfig(bool(enabled), recipe, FP8Format(fp8_format)) self._config = _ActiveConfig(bool(enabled), recipe, FP8Format(fp8_format))
self._tokens: List[Token] = [] self._tokens: List[Token] = []
@@ -322,17 +287,17 @@ def fp8_linear_forward(
Composed from the two stateless primitives: quantize x/w with the active Composed from the two stateless primitives: quantize x/w with the active
scales, run the pre-quantized GEMM with the bias fused into its epilogue. scales, run the pre-quantized GEMM with the bias fused into its epilogue.
Delayed scaling folds Delayed scaling lets the quantize kernel fold the fused amax into the
the returned amax into the history ring and publishes the next scale; history ring and publish the next scale in its own last block; dynamic
dynamic scaling measures the current amax itself. Training quantizes the scaling measures the current amax itself. Training quantizes the weight
weight every step (the optimizer bumps its version, so there is no cast every step (the optimizer bumps its version, so there is no cast cache,
cache, matching ``cached_cast``-less behavior). matching ``cached_cast``-less behavior).
""" """
state = fp8_state() state = fp8_state()
if cfg is None: if cfg is None:
cfg = _current_config() cfg = _current_config()
fmt = cfg.fp8_format.fwd() fmt = cfg.fp8_format.fwd()
if isinstance(cfg.recipe, DynamicScaling): if cfg.recipe.dynamic:
sx = _dynamic_scale(x.reshape(-1, w.size(1)), cfg.recipe, fmt) sx = _dynamic_scale(x.reshape(-1, w.size(1)), cfg.recipe, fmt)
sw = _dynamic_scale(w, cfg.recipe, fmt) sw = _dynamic_scale(w, cfg.recipe, fmt)
x8, _ = quantize(x, sx.reciprocal(), fmt) x8, _ = quantize(x, sx.reciprocal(), fmt)
@@ -345,25 +310,25 @@ def fp8_linear_forward(
).reshape(*x.shape[:-1], w.size(0)) ).reshape(*x.shape[:-1], w.size(0))
return out, sx, sw return out, sx, sw
meta = state.get_weight_meta(w) meta = state.get_weight_meta(w, cfg.recipe)
if not meta.w.initialized: if not meta.w.initialized:
meta.w.seed(w, fmt) meta.w.seed(w, fmt)
if not meta.x.initialized: if not meta.x.initialized:
meta.x.seed(x, fmt) meta.x.seed(x, fmt)
sx, sw = meta.x.scale.clone(), meta.w.scale.clone() sx, sw = meta.x.scale.clone(), meta.w.scale.clone()
x8, amax_x = quantize(x, sx.reciprocal(), fmt) # The clones feed this call's kernels (stream-ordered before the in-kernel
# fold overwrites the ring scale slots); the fp8 quantize kernel folds the
# amax into the history window and publishes the next scale itself.
x8, _ = quantize(x, sx.reciprocal(), fmt, **meta.x.fold_args(fmt))
if _is_fp8(w.dtype): if _is_fp8(w.dtype):
w8, amax_w = w, None w8 = w
else: else:
w8, amax_w = quantize(w, sw.reciprocal(), fmt) w8, _ = quantize(w, sw.reciprocal(), fmt, **meta.w.fold_args(fmt))
out = mm_fp8( out = mm_fp8(
x8.reshape(-1, x8.size(-1)), w8, sx * sw, trans_b=True, bias=bias x8.reshape(-1, x8.size(-1)), w8, sx * sw, trans_b=True, bias=bias
).reshape(*x.shape[:-1], w.size(0)) ).reshape(*x.shape[:-1], w.size(0))
meta.x.update(amax_x, fmt)
if amax_w is not None:
meta.w.update(amax_w, fmt)
meta.x.advance() meta.x.advance()
if amax_w is not None: if not _is_fp8(w.dtype):
meta.w.advance() meta.w.advance()
return out, sx, sw return out, sx, sw
@@ -385,8 +350,8 @@ class _LinearFp8(torch.autograd.Function):
ctx.save_for_backward(x, w, sx, sw) ctx.save_for_backward(x, w, sx, sw)
ctx.fmt_bwd = cfg.fp8_format.bwd() ctx.fmt_bwd = cfg.fp8_format.bwd()
ctx.recipe = cfg.recipe ctx.recipe = cfg.recipe
ctx.is_dynamic = isinstance(cfg.recipe, DynamicScaling) ctx.is_dynamic = cfg.recipe.dynamic
ctx.meta = None if ctx.is_dynamic else _state.get_weight_meta(w) ctx.meta = None if ctx.is_dynamic else _state.get_weight_meta(w, cfg.recipe)
return out return out
@staticmethod @staticmethod
@@ -411,22 +376,26 @@ class _LinearFp8(torch.autograd.Function):
# quantize outputs: g8 [m,n] with w8T [k,n] (trans_b=True) gives # quantize outputs: g8 [m,n] with w8T [k,n] (trans_b=True) gives
# grad_x, g8T [n,m] with x8T [k,m] gives grad_w — no NN-swap or TT # grad_x, g8T [n,m] with x8T [k,m] gives grad_w — no NN-swap or TT
# crosswise kernel in the training path. g is consumed in both # crosswise kernel in the training path. g is consumed in both
# orientations, so one dual-layout pass feeds both. # orientations, so quantize_dual's single pass feeds both.
g8, g8T, amax_g = quantize(g2, sg.reciprocal(), fmt, layout=2) # The g quantize folds the gradient amax into its ring in-kernel;
x8T, _ = quantize(x.reshape(-1, x.size(-1)), sx.reciprocal(), fmt, layout=1) # the x8T/w8T orientation copies discard amax (those rings were
# folded at forward time).
g8, g8T, _ = quantize_dual(g2, sg.reciprocal(), fmt, **meta.g.fold_args(fmt))
x8T, _ = quantize(
x.reshape(-1, x.size(-1)), sx.reciprocal(), fmt, transposed=True
)
if _is_fp8(w.dtype): if _is_fp8(w.dtype):
# Pre-quantized weight has no transposed copy: keep the swap # Pre-quantized weight has no transposed copy: keep the swap
# path for grad_x (grad_w is unaffected). # path for grad_x (grad_w is unaffected).
grad_x = mm_fp8(g8, w, sg * sw).reshape(x.shape) grad_x = mm_fp8(g8, w, sg * sw).reshape(x.shape)
else: else:
w8T, _ = quantize(w, sw.reciprocal(), fmt, layout=1) w8T, _ = quantize(w, sw.reciprocal(), fmt, transposed=True)
grad_x = mm_fp8(g8, w8T, sg * sw, trans_b=True).reshape(x.shape) grad_x = mm_fp8(g8, w8T, sg * sw, trans_b=True).reshape(x.shape)
grad_w = mm_fp8(g8T, x8T, sg * sx, trans_b=True) # g8.T @ x8 grad_w = mm_fp8(g8T, x8T, sg * sx, trans_b=True) # g8.T @ x8
# bias-free linears must not pay the column-sum # bias-free linears must not pay the column-sum
# reduce: g2.sum(0) is another full read of the gradient. # reduce: g2.sum(0) is another full read of the gradient.
grad_b = g2.sum(0).to(torch.bfloat16) if ctx.needs_input_grad[2] else None grad_b = g2.sum(0).to(torch.bfloat16) if ctx.needs_input_grad[2] else None
if not ctx.is_dynamic: if not ctx.is_dynamic:
meta.g.update(amax_g, fmt)
meta.g.advance() meta.g.advance()
return grad_x, grad_w, grad_b return grad_x, grad_w, grad_b
+65 -200
View File
@@ -1,14 +1,20 @@
"""FP8 CUDA kernel interface adapter (the only module touching the pybind). """FP8 CUDA kernel interface adapter (the only module touching the pybind).
Isolates the ``fp8_ops`` CUDA extension behind stable Python primitives: Attention-style thin wrappers: one Python entry per binding, called directly
— no torch.library dispatch layer. Optional arguments (``ring_state``,
``bias``) keep native Optional semantics at the pybind boundary, and
in-place buffer updates (the delayed-scaling ring fold, like attention's
KV-cache appends) happen on-stream without mutation declarations. CUDA-only:
non-CUDA or unsupported inputs raise from the binding's TORCH_CHECKs.
- ``quantize(x, scale, fmt) -> (x8, amax)`` — BF16/FP16/FP32 → FP8 with fused amax - ``quantize(x, scale, fmt, transposed=False) -> (x8|x8T, amax)`` — BF16/FP16/FP32
→ FP8 with fused amax (``transposed`` picks the orientation; arity is fixed)
- ``quantize_dual(x, scale, fmt) -> (x8, x8T, amax)`` — both orientations, one read
- ``mm_fp8(a8, b8, sa, sb) -> out`` — pre-quantized FP8 GEMM (BF16 output) - ``mm_fp8(a8, b8, sa, sb) -> out`` — pre-quantized FP8 GEMM (BF16 output)
Scale semantics: scales are *quantization steps* — the value divided out when ``scale`` is the quantization multiplier (device scalar); ``fmt`` is
quantizing (``x8 = x / scale``). Every primitive computes its own inverse ``"e4m3"`` or ``"e5m2"``. ``amax`` values are *returned*, never passed as
internally; callers never pass ``scale_inv``. ``amax`` values are *returned*, output arguments.
never passed as output arguments. ``fmt`` is ``"e4m3"`` or ``"e5m2"``.
Policy (scales, amax history, delayed scaling, autocast) lives in ``fp8.py``; Policy (scales, amax history, delayed scaling, autocast) lives in ``fp8.py``;
this module is stateless. this module is stateless.
@@ -17,7 +23,6 @@ this module is stateless.
from typing import Optional, Tuple from typing import Optional, Tuple
import torch import torch
from torch.library import custom_op
from astrai.extension.loader import get_module from astrai.extension.loader import get_module
@@ -32,192 +37,63 @@ def _fmt_int(fmt: str) -> int:
raise ValueError(f"unsupported fp8 format {fmt!r} (expected 'e4m3' or 'e5m2')") raise ValueError(f"unsupported fp8 format {fmt!r} (expected 'e4m3' or 'e5m2')")
def _fmt_name(fmt: int) -> str:
if fmt == 0:
return "e4m3"
if fmt == 1:
return "e5m2"
raise ValueError(f"unsupported quantization type {fmt!r}")
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]:
"""Float (bf16/fp16/fp32) -> FP8 quantize with fused amax; ``scale`` is a multiplier."""
@fp8_quantize.register_fake
def _fp8_quantize_fake(x, scale, fmt):
dtype = torch.float8_e5m2 if fmt == 1 else torch.float8_e4m3fn
return (
torch.empty(x.shape, device=x.device, dtype=dtype),
torch.empty(1, device=x.device, dtype=torch.float32),
)
_QUANT_INPUT_DTYPES = (torch.bfloat16, torch.float16, torch.float32)
@custom_op("custom::fp8_quantize_t", mutates_args=())
def fp8_quantize_t(
x: torch.Tensor, scale: torch.Tensor, fmt: int
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Transposed-output variant of fp8_quantize: returns ``(x8T, amax)``
where ``x8T`` is the [cols][rows] row-major transpose of the quantized
input (the K-contiguous operand orientation for NT GEMMs)."""
@fp8_quantize_t.register_fake
def _fp8_quantize_t_fake(x, scale, fmt):
dtype = torch.float8_e5m2 if fmt == 1 else torch.float8_e4m3fn
rows, cols = x.shape[-2], x.shape[-1]
return (
torch.empty((*x.shape[:-2], cols, rows), device=x.device, dtype=dtype),
torch.empty(1, device=x.device, dtype=torch.float32),
)
@fp8_quantize_t.register_kernel("cuda")
def _fp8_quantize_t_cuda(x, scale, fmt):
if x.dtype not in _QUANT_INPUT_DTYPES:
raise TypeError(f"fp8 quantize requires bf16/fp16/fp32 input, got {x.dtype}")
return get_module("fp8_ops").quantize(x, scale, int(fmt), 1)
@fp8_quantize_t.register_kernel("cpu")
def _fp8_quantize_t_cpu(x, scale, fmt):
x8, amax = _fp8_quantize_cpu(x, scale, fmt)
return x8.transpose(-2, -1).contiguous(), amax
@custom_op("custom::fp8_quantize_dual", mutates_args=())
def fp8_quantize_dual(
x: torch.Tensor, scale: torch.Tensor, fmt: int
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""Dual-orientation quantize: one read of ``x`` produces both the
row-major ``x8`` and its transposed ``x8T`` (plus ``amax``), for tensors
consumed by GEMMs on both orientations (backward ``g``)."""
@fp8_quantize_dual.register_fake
def _fp8_quantize_dual_fake(x, scale, fmt):
dtype = torch.float8_e5m2 if fmt == 1 else torch.float8_e4m3fn
rows, cols = x.shape[-2], x.shape[-1]
return (
torch.empty(x.shape, device=x.device, dtype=dtype),
torch.empty((*x.shape[:-2], cols, rows), device=x.device, dtype=dtype),
torch.empty(1, device=x.device, dtype=torch.float32),
)
@fp8_quantize_dual.register_kernel("cuda")
def _fp8_quantize_dual_cuda(x, scale, fmt):
if x.dtype not in _QUANT_INPUT_DTYPES:
raise TypeError(f"fp8 quantize requires bf16/fp16/fp32 input, got {x.dtype}")
return get_module("fp8_ops").quantize(x, scale, int(fmt), 2)
@fp8_quantize_dual.register_kernel("cpu")
def _fp8_quantize_dual_cpu(x, scale, fmt):
x8, amax = _fp8_quantize_cpu(x, scale, fmt)
return x8, x8.transpose(-2, -1).contiguous(), amax
@fp8_quantize.register_kernel("cuda")
def _fp8_quantize_cuda(x, scale, fmt):
if x.dtype not in _QUANT_INPUT_DTYPES:
raise TypeError(f"fp8 quantize requires bf16/fp16/fp32 input, got {x.dtype}")
return get_module("fp8_ops").quantize(x, scale, int(fmt))
@fp8_quantize.register_kernel("cpu")
def _fp8_quantize_cpu(x, scale, fmt):
x8 = (x.float() * scale).to(_fmt_dtype(_fmt_name(fmt)))
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,
scale: torch.Tensor,
trans_a: int = 0,
trans_b: int = 0,
bias: Optional[torch.Tensor] = None,
) -> torch.Tensor:
"""FP8 GEMM: ``a @ b * scale (+ bias)`` with FP32 accumulation.
2D or 3D (batched) operands; a size-1 batch broadcasts (matmul rules).
``bias`` (bf16, length n) fuses into the epilogue in fp32 before the
single bf16 rounding. The result is always BF16; FP8 output is a
separate quantize operation.
"""
@fp8_gemm.register_fake
def _fp8_gemm_fake(a, b, scale, trans_a=0, trans_b=0, bias=None):
dtype = torch.bfloat16
rows = a.size(2) if trans_a else a.size(1)
cols = b.size(1) if trans_b else b.size(2)
batches = [t.size(0) for t in (a, b) if t.dim() == 3]
shape = (max(batches), rows, cols) if batches else (rows, cols)
return torch.empty(shape, device=a.device, dtype=dtype)
@fp8_gemm.register_kernel("cuda")
def _fp8_gemm_cuda(a, b, scale, trans_a=0, trans_b=0, bias=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, scale, trans_a, trans_b, bias)
@fp8_gemm.register_kernel("cpu")
def _fp8_gemm_cpu(a, b, scale, trans_a=0, trans_b=0, bias=None):
aa = a.float().transpose(-2, -1) if trans_a else a.float()
bb = b.float().transpose(-2, -1) if trans_b else b.float()
acc = aa @ bb * scale
if bias is not None and bias.numel() > 0:
acc = acc + bias.float()
return acc.to(torch.bfloat16)
def quantize( def quantize(
x: torch.Tensor, scale: torch.Tensor, fmt: str = "e4m3", layout: int = 0 x: torch.Tensor,
) -> tuple: scale: torch.Tensor,
fmt: str = "e4m3",
transposed: bool = False,
ring_state: Optional[torch.Tensor] = None,
hist_idx: int = 0,
fp8_max: float = 448.0,
pow2_margin: float = 1.0,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Float (bf16/fp16/fp32) -> FP8 quantize with fused amax. """Float (bf16/fp16/fp32) -> FP8 quantize with fused amax.
``scale`` is the quantization multiplier (device scalar); ``fmt`` selects ``scale`` is the quantization multiplier (device scalar); ``fmt`` selects
E4M3 or E5M2. ``amax`` is a fresh 1-element float32 tensor. ``layout`` E4M3 or E5M2. ``amax`` is a fresh 1-element float32 tensor.
picks the output orientation: 0 = row-major ``(x8, amax)``; 1 = ``transposed=True`` swaps ``x8`` for ``x8T``, the ``[cols][rows]``
transposed ``[cols][rows]`` ``(x8T, amax)`` — the K-contiguous operand row-major transpose of the quantized input — the K-contiguous operand
orientation NT GEMMs want; 2 = both from one read ``(x8, x8T, amax)`` orientation NT GEMMs want — at the same 2-tuple arity.
(for tensors consumed in both orientations, e.g. backward ``g``).
``ring_state`` (a 1D float32 CUDA buffer laid out
``[hist n | scale | legacy | amax | done]``) switches on the in-kernel
delayed-scaling fold: the kernel's last block folds the amax into
``hist[hist_idx]`` and publishes the next scale as
``max(hist) / fp8_max / pow2_margin`` — the returned ``amax`` is then the
self-cleaned persistent slot (reads zero). None keeps the classic
fresh-amax return.
""" """
# Hot-path bypass of the torch.library dispatch (~5us/call, ~40% of a return get_module("fp8_ops").quantize(
# 512-wide GEMM): real CUDA tensors of a supported dtype go straight to x,
# the extension. Fake/subclass tensors and non-CUDA inputs keep the scale,
# custom_op route so torch.compile / meta / fake-tensor tracing and the _fmt_int(fmt),
# CPU fallback behave exactly as before. transposed,
if ( ring_state,
type(x) is torch.Tensor hist_idx,
and x.is_cuda fp8_max,
and x.dtype in _QUANT_INPUT_DTYPES pow2_margin,
and fmt in _FMT_TO_INT )
):
return get_module("fp8_ops").quantize(x, scale, _FMT_TO_INT[fmt], layout)
if layout == 0: def quantize_dual(
return fp8_quantize(x, scale, _fmt_int(fmt)) x: torch.Tensor,
if layout == 1: scale: torch.Tensor,
return fp8_quantize_t(x, scale, _fmt_int(fmt)) fmt: str = "e4m3",
return fp8_quantize_dual(x, scale, _fmt_int(fmt)) ring_state: Optional[torch.Tensor] = None,
hist_idx: int = 0,
fp8_max: float = 448.0,
pow2_margin: float = 1.0,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""Dual-orientation quantize: one read of ``x`` produces both the
row-major ``x8`` and its transposed ``x8T`` (plus ``amax``), for tensors
consumed by GEMMs in both orientations (backward ``g``).
``ring_state`` switches on the in-kernel delayed-scaling fold exactly as
in :func:`quantize`.
"""
return get_module("fp8_ops").quantize_dual(
x, scale, _fmt_int(fmt), ring_state, hist_idx, fp8_max, pow2_margin
)
def mm_fp8( def mm_fp8(
@@ -237,15 +113,4 @@ def mm_fp8(
kernel epilogue in fp32 — no separate elementwise pass. The result is kernel epilogue in fp32 — no separate elementwise pass. The result is
BF16; FP8 output is a separate quantize operation. BF16; FP8 output is a separate quantize operation.
""" """
# Same hot-path bypass as quantize(): the binding's TORCH_CHECKs keep return get_module("fp8_ops").mm_fp8(a, b, scale, trans_a, trans_b, bias)
# validation identical on the direct route (bias may be None — the
# binding resolves it to the no-bias path).
if (
type(a) is torch.Tensor
and a.is_cuda
and a.dtype in (torch.float8_e4m3fn, torch.float8_e5m2)
):
return get_module("fp8_ops").mm_fp8(
a, b, scale, int(trans_a), int(trans_b), bias
)
return fp8_gemm(a, b, scale, trans_a, trans_b, bias)
+38 -16
View File
@@ -67,11 +67,15 @@ class DecodeSteadyState:
When the same ordered task set decodes one token per step, sampling When the same ordered task set decodes one token per step, sampling
params and task signature are reused; only positions advance by 1. params and task signature are reused; only positions advance by 1.
``last_tokens`` keeps that step's sampled ids on-device so the next
step with an unchanged signature can fill ``input_ids`` via a
device-to-device copy.
""" """
task_sig: tuple task_sig: tuple
positions: list[int] positions: list[int]
sampling_info: SamplingBatchInfo sampling_info: SamplingBatchInfo
last_tokens: Optional[Tensor] = None
def _build_sampling_batch_info(tasks: List[Task], device) -> SamplingBatchInfo: def _build_sampling_batch_info(tasks: List[Task], device) -> SamplingBatchInfo:
@@ -250,6 +254,13 @@ class Executor:
return_logprobs: bool = False, return_logprobs: bool = False,
info: Optional[SamplingBatchInfo] = None, info: Optional[SamplingBatchInfo] = None,
): ):
"""Sample from ``logits`` and return ``(host_payload, tokens)``.
``host_payload`` is the scheduler-facing list (token ids, or
``(token_id, logprob)`` tuples with ``return_logprobs``);
``tokens`` is the ``[B]`` device tensor that produced it, kept
for the steady-state decode fast path.
"""
info = info or _build_sampling_batch_info(tasks, self.device) info = info or _build_sampling_batch_info(tasks, self.device)
if info.has_freq: if info.has_freq:
history_lists = [ history_lists = [
@@ -284,14 +295,14 @@ class Executor:
return_logprobs=return_logprobs, return_logprobs=return_logprobs,
) )
if not return_logprobs: if not return_logprobs:
return result.tolist() return result.tolist(), result
tokens, logprobs = result tokens, logprobs = result
tokens_list = tokens.tolist() tokens_list = tokens.tolist()
logprobs_list = logprobs.tolist() logprobs_list = logprobs.tolist()
for task, logprob in zip(tasks, logprobs_list): for task, logprob in zip(tasks, logprobs_list):
task.output_logprobs.append(float(logprob)) task.output_logprobs.append(float(logprob))
return list(zip(tokens_list, logprobs_list)) return list(zip(tokens_list, logprobs_list)), tokens
def execute_prefill( def execute_prefill(
self, self,
@@ -336,7 +347,8 @@ class Executor:
torch.arange(1, batch_sz + 1, device=self.device) * q_len - 1 torch.arange(1, batch_sz + 1, device=self.device) * q_len - 1
] ]
return tasks, self._sample_logits(logits, tasks, return_logprobs) step_out, _ = self._sample_logits(logits, tasks, return_logprobs)
return tasks, step_out
def execute_decode( def execute_decode(
self, tasks: List[Task], return_logprobs: bool = False self, tasks: List[Task], return_logprobs: bool = False
@@ -360,24 +372,30 @@ class Executor:
b = len(tasks) b = len(tasks)
ws = self._workspace ws = self._workspace
task_ids = [t.task_id for t in tasks]
cur_positions = [t.next_pos for t in tasks]
task_sig = tuple(task_ids)
# ---- pre-replay: update input buffers in-place ---- # ---- pre-replay: update input buffers in-place ----
input_ids = ws.fill_input_ids( # When the previous decode step ran this same ordered task set, its
[t.output_ids[-1] if t.output_ids else t.prompt_ids[-1] for t in tasks] # sampled tokens are still on-device and map 1:1 onto the current
) # slots — fill input ids device-to-device. inference_mode guards
# the read because the source was produced under sampling's
task_ids = [t.task_id for t in tasks] # inference-mode context.
cur_positions = [t.next_pos for t in tasks] cached = self._decode_cache
sig_match = cached is not None and cached.task_sig == task_sig
if sig_match and cached.last_tokens is not None:
with torch.inference_mode():
input_ids = ws.fill_input_ids_from_device(cached.last_tokens)
else:
input_ids = ws.fill_input_ids(
[t.output_ids[-1] if t.output_ids else t.prompt_ids[-1] for t in tasks]
)
kv_cache = self.task_cache.bind(task_ids, ws) kv_cache = self.task_cache.bind(task_ids, ws)
task_sig = tuple(task_ids) reuse_decode_state = self.task_cache.bind_was_steady and sig_match
reuse_decode_state = (
self.task_cache.bind_was_steady
and self._decode_cache is not None
and self._decode_cache.task_sig == task_sig
)
if reuse_decode_state: if reuse_decode_state:
info = self._decode_cache.sampling_info info = self._decode_cache.sampling_info
ws.position_ids[:b] += 1 ws.position_ids[:b] += 1
@@ -418,4 +436,8 @@ class Executor:
) )
logits = outputs["logits"] logits = outputs["logits"]
return self._sample_logits(logits, tasks, return_logprobs, info=info) step_out, tokens_dev = self._sample_logits(
logits, tasks, return_logprobs, info=info
)
self._decode_cache.last_tokens = tokens_dev
return step_out
+12
View File
@@ -139,6 +139,18 @@ class InferenceWorkspace:
self.input_ids[:b].copy_(pin[:b]) self.input_ids[:b].copy_(pin[:b])
return self.input_ids[:b] return self.input_ids[:b]
def fill_input_ids_from_device(self, tokens: Tensor) -> Tensor:
"""Copy device-resident ``[B]`` token ids into the device buffer.
Steady-state decode fast path: when the executor's cached task
signature still matches, the previous step's sampled tokens map
1:1 onto the current slots, so the ids transfer device-to-device
instead of round-tripping through the host staging buffers.
"""
b = tokens.size(0)
self.input_ids[:b].copy_(tokens)
return self.input_ids[:b]
def decode_mask(self, position_ids: Tensor, total_len: int) -> Tensor: def decode_mask(self, position_ids: Tensor, total_len: int) -> Tensor:
"""Return the ``[B, 1, total_len]`` validity mask for this step. """Return the ``[B, 1, total_len]`` validity mask for this step.
+2
View File
@@ -116,6 +116,8 @@ class GradientCheckpointingCallback(TrainCallback):
del module._original_forward del module._original_forward
def on_train_begin(self, context: TrainContext): def on_train_begin(self, context: TrainContext):
if not self.modules:
return
context.model.apply(self._enable) context.model.apply(self._enable)
logger.info("Gradient checkpointing enabled") logger.info("Gradient checkpointing enabled")
+25 -4
View File
@@ -54,19 +54,40 @@ struct Fp8GemmTraits {
"warp tile must be a multiple of the m16n8 MMA shape"); "warp tile must be a multiple of the m16n8 MMA shape");
}; };
// Quantize output orientation: RowMajor = x8 only; Transposed = the
// [cols][rows] x8T only; Dual = both from a single read. Transposed/Dual
// produce K-contiguous operands so crosswise consumers (backward
// grad_x / grad_w) route through the NT fast path.
enum class QuantLayout : int {
RowMajor = 0,
Transposed = 1,
Dual = 2,
};
// Quantize-kernel parameter POD: float input -> FP8 with fused amax. // Quantize-kernel parameter POD: float input -> FP8 with fused amax.
struct FP8QuantizeParams { struct FP8QuantizeParams {
const void* __restrict__ input_ptr = nullptr; const void* __restrict__ input_ptr = nullptr;
void* __restrict__ output_ptr = nullptr; void* __restrict__ output_ptr = nullptr;
void* __restrict__ output_transposed_ptr = nullptr; // [cols][rows] void* __restrict__ output_transposed_ptr = nullptr; // [cols][rows]
// Output layout: 0 = row-major only, 1 = transposed only, 2 = both from QuantLayout out_layout = QuantLayout::RowMajor;
// a single read. Modes 1/2 produce K-contiguous operands so crosswise
// consumers (backward grad_x / grad_w) route through the NT fast path.
int out_layout = 0;
const float* __restrict__ scale = nullptr; // device multiplier const float* __restrict__ scale = nullptr; // device multiplier
float* __restrict__ amax = nullptr; // raw-domain max out float* __restrict__ amax = nullptr; // raw-domain max out
// Optional delayed-scaling ring fold: when fold_ring is set, the kernel's
// last-finishing block folds the final amax into hist[hist_idx], reduces
// the window and publishes the next scale — replacing the host-side
// update chain. amax then points at a persistent self-cleaning slot
// (zeroed by the same last block) inside the caller's ring state.
bool fold_ring = false;
float* __restrict__ hist = nullptr; // [hist_len] amax history window
float* __restrict__ scale_out = nullptr;
unsigned int* __restrict__ done = nullptr; // block-completion counter
int hist_len = 0;
int hist_idx = 0;
float fp8_max = 448.0f; // scale = max(hist) / fp8_max / pow2_margin
float pow2_margin = 1.0f;
// Element count (elementwise kernel); the tiled kernel views the same // Element count (elementwise kernel); the tiled kernel views the same
// buffer as [rows][cols] row-major. // buffer as [rows][cols] row-major.
int total = 0; int total = 0;
+109 -52
View File
@@ -1,4 +1,4 @@
// CUDA bindings for the two stateless FP8 primitives. // CUDA bindings for the stateless FP8 quantize/GEMM primitives.
#include <ATen/cuda/CUDAContext.h> #include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.h> #include <c10/cuda/CUDAGuard.h>
@@ -94,13 +94,17 @@ void launch_quantize_for(const torch::Tensor& x, const FP8QuantizeParams& p,
launch_for_dtype<Tiled, FP8Format::E4M3>(x, p, stream); launch_for_dtype<Tiled, FP8Format::E4M3>(x, p, stream);
} }
} // namespace // Shared binding body for the two quantize entry points: RowMajor /
// Transposed (single output) serve quantize(), Dual (both orientations from
// Output-layout dispatch: 0 = [rows][cols] row-major (2-tuple return), // one read) serves quantize_dual(). A ring tensor switches
// 1 = transposed [cols][rows] only (2-tuple), 2 = both orientations from a // on the in-kernel delayed-scaling fold: state layout
// single read (3-tuple). Layouts 1/2 feed the NT GEMM fast path. // [hist n | scale | legacy | amax | done-as-int], and the returned amax is
py::object quantize(torch::Tensor x, torch::Tensor scale, int64_t fmt, // the (self-cleaned) persistent slot. Without it, amax is reduced into a
int64_t layout) { // fresh buffer armed by a driver memset — cheaper than the zeros() fill
// kernel.
py::object quantize_impl(torch::Tensor x, torch::Tensor scale, int64_t fmt,
QuantLayout layout, py::object ring, int64_t hist_idx,
double fp8_max, double pow2_margin) {
TORCH_CHECK(x.is_cuda(), "CUDA tensors required"); TORCH_CHECK(x.is_cuda(), "CUDA tensors required");
TORCH_CHECK(x.scalar_type() == torch::kBFloat16 || TORCH_CHECK(x.scalar_type() == torch::kBFloat16 ||
x.scalar_type() == torch::kHalf || x.scalar_type() == torch::kHalf ||
@@ -109,9 +113,7 @@ py::object quantize(torch::Tensor x, torch::Tensor scale, int64_t fmt,
TORCH_CHECK(fmt == static_cast<int64_t>(FP8Format::E4M3) || TORCH_CHECK(fmt == static_cast<int64_t>(FP8Format::E4M3) ||
fmt == static_cast<int64_t>(FP8Format::E5M2), fmt == static_cast<int64_t>(FP8Format::E5M2),
"unsupported quantization type: expected E4M3 (0) or E5M2 (1)"); "unsupported quantization type: expected E4M3 (0) or E5M2 (1)");
TORCH_CHECK(layout >= 0 && layout <= 2, TORCH_CHECK(layout == QuantLayout::RowMajor || x.dim() >= 2,
"layout must be 0 (row-major), 1 (transposed) or 2 (both)");
TORCH_CHECK(layout == 0 || x.dim() >= 2,
"transposed quantize layouts need a 2D+ tensor"); "transposed quantize layouts need a 2D+ tensor");
check_scale(scale, x); check_scale(scale, x);
check_fp8_device(x); check_fp8_device(x);
@@ -120,37 +122,94 @@ py::object quantize(torch::Tensor x, torch::Tensor scale, int64_t fmt,
auto input = x.contiguous(); auto input = x.contiguous();
auto out_opts = input.options().dtype( auto out_opts = input.options().dtype(
fmt ? torch::kFloat8_e5m2 : torch::kFloat8_e4m3fn); fmt ? torch::kFloat8_e5m2 : torch::kFloat8_e4m3fn);
auto amax = torch::zeros({1}, input.options().dtype(torch::kFloat32)); torch::Tensor amax;
float *ring_hist = nullptr, *ring_scale_out = nullptr;
unsigned int* ring_done = nullptr;
int ring_len = 0;
if (!ring.is_none()) {
auto st = ring.cast<torch::Tensor>();
TORCH_CHECK(st.is_cuda() && st.dim() == 1 &&
st.scalar_type() == torch::kFloat32,
"ring state must be a 1D float32 CUDA tensor");
const int64_t n = st.numel() - 4;
TORCH_CHECK(n > 0 && hist_idx >= 0 && hist_idx < n,
"ring state too small or hist_idx out of range");
float* base = st.data_ptr<float>();
amax = st.narrow(0, n + 2, 1);
ring_hist = base;
ring_scale_out = base + n;
ring_done = reinterpret_cast<unsigned int*>(base + n + 3);
ring_len = static_cast<int>(n);
} else {
amax = torch::empty({1}, input.options().dtype(torch::kFloat32));
cudaMemsetAsync(amax.data_ptr(), 0, sizeof(float), stream.stream());
}
FP8QuantizeParams p; FP8QuantizeParams p;
p.input_ptr = input.data_ptr(); p.input_ptr = input.data_ptr();
p.scale = scale.data_ptr<float>(); p.scale = scale.data_ptr<float>();
p.amax = amax.data_ptr<float>(); p.amax = amax.data_ptr<float>();
if (ring_hist) {
p.fold_ring = true;
p.hist = ring_hist;
p.scale_out = ring_scale_out;
p.done = ring_done;
p.hist_len = ring_len;
p.hist_idx = static_cast<int>(hist_idx);
p.fp8_max = static_cast<float>(fp8_max);
p.pow2_margin = static_cast<float>(pow2_margin);
}
p.total = static_cast<int>(input.numel()); p.total = static_cast<int>(input.numel());
p.out_layout = static_cast<int>(layout); p.out_layout = layout;
p.rows = static_cast<int>(input.size(-2)); p.rows = static_cast<int>(input.size(-2));
p.cols = static_cast<int>(input.size(-1)); p.cols = static_cast<int>(input.size(-1));
torch::Tensor output, output_t; torch::Tensor output, output_t;
if (layout == 0 || layout == 2) { if (layout != QuantLayout::Transposed) {
output = torch::empty_like(input, out_opts); output = torch::empty_like(input, out_opts);
p.output_ptr = output.data_ptr(); p.output_ptr = output.data_ptr();
} }
if (layout >= 1) { if (layout != QuantLayout::RowMajor) {
output_t = torch::empty({input.size(-1), input.size(-2)}, out_opts); output_t = torch::empty({input.size(-1), input.size(-2)}, out_opts);
p.output_transposed_ptr = output_t.data_ptr(); p.output_transposed_ptr = output_t.data_ptr();
} }
const bool e5m2 = fmt == static_cast<int64_t>(FP8Format::E5M2); const bool e5m2 = fmt == static_cast<int64_t>(FP8Format::E5M2);
if (layout != 0) if (layout == QuantLayout::RowMajor)
launch_quantize_for<true>(input, p, e5m2, stream.stream());
else
launch_quantize_for<false>(input, p, e5m2, stream.stream()); launch_quantize_for<false>(input, p, e5m2, stream.stream());
else
launch_quantize_for<true>(input, p, e5m2, stream.stream());
C10_CUDA_CHECK(cudaGetLastError()); C10_CUDA_CHECK(cudaGetLastError());
if (layout == 2) return py::make_tuple(output, output_t, amax); if (layout == QuantLayout::Dual)
return py::make_tuple(layout == 1 ? output_t : output, amax); return py::make_tuple(output, output_t, amax);
return py::make_tuple(
layout == QuantLayout::Transposed ? output_t : output, amax);
}
} // namespace
// Single-orientation quantize binding: row-major x8, or its [cols][rows]
// transpose when transposed is set — the K-contiguous operand orientation
// NT GEMMs want. Returns (x8|x8T, amax).
py::object quantize(torch::Tensor x, torch::Tensor scale, int64_t fmt,
bool transposed, py::object ring, int64_t hist_idx,
double fp8_max, double pow2_margin) {
const QuantLayout layout =
transposed ? QuantLayout::Transposed : QuantLayout::RowMajor;
return quantize_impl(x, scale, fmt, layout, ring, hist_idx, fp8_max,
pow2_margin);
}
// Dual-orientation quantize binding: one read of x produces both the
// row-major x8 and its transpose (plus amax), for tensors consumed by GEMMs
// in both orientations (backward g). Returns (x8, x8T, amax).
py::object quantize_dual(torch::Tensor x, torch::Tensor scale, int64_t fmt,
py::object ring, int64_t hist_idx, double fp8_max,
double pow2_margin) {
return quantize_impl(x, scale, fmt, QuantLayout::Dual, ring, hist_idx,
fp8_max, pow2_margin);
} }
torch::Tensor mm_fp8(torch::Tensor a, torch::Tensor b, torch::Tensor scale, torch::Tensor mm_fp8(torch::Tensor a, torch::Tensor b, torch::Tensor scale,
int64_t trans_a, int64_t trans_b, torch::Tensor bias) { bool trans_a, bool trans_b, py::object bias) {
TORCH_CHECK(a.is_cuda() && b.is_cuda(), "CUDA tensors required"); TORCH_CHECK(a.is_cuda() && b.is_cuda(), "CUDA tensors required");
TORCH_CHECK(a.scalar_type() == torch::kFloat8_e4m3fn || TORCH_CHECK(a.scalar_type() == torch::kFloat8_e4m3fn ||
a.scalar_type() == torch::kFloat8_e5m2, a.scalar_type() == torch::kFloat8_e5m2,
@@ -160,6 +219,18 @@ torch::Tensor mm_fp8(torch::Tensor a, torch::Tensor b, torch::Tensor scale,
(b.dim() == 2 || b.dim() == 3), (b.dim() == 2 || b.dim() == 3),
"a and b must be 2D or 3D (batched)"); "a and b must be 2D or 3D (batched)");
TORCH_CHECK(a.device() == b.device(), "a and b must share device"); TORCH_CHECK(a.device() == b.device(), "a and b must share device");
// Python None and an omitted argument both mean "no bias" — an undefined
// tensor below. (py::isinstance<torch::Tensor> is false for real tensors
// here — torch's caster registers no pybind type info — so validate by
// attempting the cast itself.)
torch::Tensor bias_t;
if (!bias.is_none()) {
try {
bias_t = bias.cast<torch::Tensor>();
} catch (const py::cast_error&) {
TORCH_CHECK(false, "bias must be a torch.Tensor or None");
}
}
check_scale(scale, a); check_scale(scale, a);
check_fp8_device(a); check_fp8_device(a);
const at::cuda::OptionalCUDAGuard guard(a.device()); const at::cuda::OptionalCUDAGuard guard(a.device());
@@ -176,10 +247,8 @@ torch::Tensor mm_fp8(torch::Tensor a, torch::Tensor b, torch::Tensor scale,
torch::Tensor a_st, b_st; torch::Tensor a_st, b_st;
int64_t a_ld, b_ld, a_bstride, b_bstride; int64_t a_ld, b_ld, a_bstride, b_bstride;
const bool tag_a = const bool tag_a = resolve_operand(a, trans_a, a_ld, a_bstride, a_st);
resolve_operand(a, trans_a != 0, a_ld, a_bstride, a_st); const bool tag_b = resolve_operand(b, trans_b, b_ld, b_bstride, b_st);
const bool tag_b =
resolve_operand(b, trans_b != 0, b_ld, b_bstride, b_st);
// GEMM dims from the user flags; storage layout never swaps them. // GEMM dims from the user flags; storage layout never swaps them.
const int64_t m = trans_a ? a.size(-1) : a.size(-2); const int64_t m = trans_a ? a.size(-1) : a.size(-2);
const int64_t k = trans_a ? a.size(-2) : a.size(-1); const int64_t k = trans_a ? a.size(-2) : a.size(-1);
@@ -203,13 +272,13 @@ torch::Tensor mm_fp8(torch::Tensor a, torch::Tensor b, torch::Tensor scale,
p.b_ld = static_cast<int>(b_ld); p.b_ld = static_cast<int>(b_ld);
// Fused epilogue bias (bf16, broadcast over rows and batches). An // Fused epilogue bias (bf16, broadcast over rows and batches). An
// undefined or 0-element tensor keeps the plain scaled output. // undefined or 0-element tensor keeps the plain scaled output.
if (bias.defined() && bias.numel() > 0) { if (bias_t.defined() && bias_t.numel() > 0) {
TORCH_CHECK(bias.is_cuda() && bias.scalar_type() == torch::kBFloat16, TORCH_CHECK(bias_t.is_cuda() && bias_t.scalar_type() == torch::kBFloat16,
"fp8 gemm bias must be a CUDA bf16 tensor"); "fp8 gemm bias must be a CUDA bf16 tensor");
TORCH_CHECK(bias.dim() == 1 && bias.size(0) == n, TORCH_CHECK(bias_t.dim() == 1 && bias_t.size(0) == n,
"fp8 gemm bias must be 1D of length n=", n); "fp8 gemm bias must be 1D of length n=", n);
TORCH_CHECK(bias.is_contiguous(), "fp8 gemm bias must be contiguous"); TORCH_CHECK(bias_t.is_contiguous(), "fp8 gemm bias must be contiguous");
p.bias_ptr = bias.data_ptr(); p.bias_ptr = bias_t.data_ptr();
} }
p.batch = static_cast<int>(batch); p.batch = static_cast<int>(batch);
p.a_batch_stride = (batch_a == 1 && batch > 1) ? 0 : a_bstride; p.a_batch_stride = (batch_a == 1 && batch > 1) ? 0 : a_bstride;
@@ -223,28 +292,16 @@ torch::Tensor mm_fp8(torch::Tensor a, torch::Tensor b, torch::Tensor scale,
return output; return output;
} }
// mm_fp8 binding: Python None and an omitted argument both mean "no bias",
// so every Python layer can pass its bias argument through untouched.
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("quantize", &quantize, py::arg("x"), py::arg("scale"), m.def("quantize", &quantize, py::arg("x"), py::arg("scale"),
py::arg("fmt"), py::arg("layout") = 0); py::arg("fmt"), py::arg("transposed") = false,
m.def( py::arg("ring") = py::none(), py::arg("hist_idx") = 0,
"mm_fp8", py::arg("fp8_max") = 448.0, py::arg("pow2_margin") = 1.0);
[](torch::Tensor a, torch::Tensor b, torch::Tensor scale, m.def("quantize_dual", &quantize_dual, py::arg("x"), py::arg("scale"),
int64_t trans_a, int64_t trans_b, py::object bias) { py::arg("fmt"), py::arg("ring") = py::none(),
torch::Tensor t; py::arg("hist_idx") = 0, py::arg("fp8_max") = 448.0,
if (!bias.is_none()) { py::arg("pow2_margin") = 1.0);
// (py::isinstance<torch::Tensor> is false for real tensors m.def("mm_fp8", &mm_fp8, py::arg("a"), py::arg("b"), py::arg("scale"),
// here — torch's caster registers no pybind type info — so py::arg("trans_a") = false, py::arg("trans_b") = false,
// validate by attempting the cast itself.) py::arg("bias") = py::none());
try {
t = bias.cast<torch::Tensor>();
} catch (const py::cast_error&) {
TORCH_CHECK(false, "bias must be a torch.Tensor or None");
}
}
return mm_fp8(a, b, scale, trans_a, trans_b, t);
},
py::arg("a"), py::arg("b"), py::arg("scale"), py::arg("trans_a") = 0,
py::arg("trans_b") = 0, py::arg("bias") = py::none());
} }
+118 -46
View File
@@ -15,8 +15,8 @@
namespace astrai { namespace astrai {
namespace fp8 { namespace fp8 {
// Input element type traits: one element -> float, and the unpack of one // Input element type traits: one element -> float, the unpack of one
// 16-byte load into kVecElems floats. // 16-byte load into kVecElems floats, and a native 2-element pair load.
template <typename InT> template <typename InT>
struct quant_in_traits; struct quant_in_traits;
@@ -37,6 +37,13 @@ struct quant_in_traits<__nv_bfloat16> {
f[2 * j + 1] = p.y; f[2 * j + 1] = p.y;
} }
} }
static __device__ __forceinline__ void load_pair(const __nv_bfloat16* p,
float* f) {
const float2 v = __bfloat1622float2(
*reinterpret_cast<const __nv_bfloat162*>(p));
f[0] = v.x;
f[1] = v.y;
}
}; };
template <> template <>
@@ -55,6 +62,13 @@ struct quant_in_traits<__half> {
f[2 * j + 1] = p.y; f[2 * j + 1] = p.y;
} }
} }
static __device__ __forceinline__ void load_pair(const __half* p,
float* f) {
const float2 v =
__half22float2(*reinterpret_cast<const __half2*>(p));
f[0] = v.x;
f[1] = v.y;
}
}; };
template <> template <>
@@ -67,6 +81,11 @@ struct quant_in_traits<float> {
#pragma unroll #pragma unroll
for (int j = 0; j < 4; ++j) f[j] = __uint_as_float(w[j]); for (int j = 0; j < 4; ++j) f[j] = __uint_as_float(w[j]);
} }
static __device__ __forceinline__ void load_pair(const float* p,
float* f) {
f[0] = p[0];
f[1] = p[1];
}
}; };
// One float -> one fp8 byte (round-nearest-even + satfinite). // One float -> one fp8 byte (round-nearest-even + satfinite).
@@ -89,8 +108,13 @@ __device__ __forceinline__ unsigned cvt_fp8x2(float a, float b) {
// Block-wide amax reduce -> one atomic per block: warp-reduce, park one // Block-wide amax reduce -> one atomic per block: warp-reduce, park one
// value per warp, thread 0 folds. kWarps must cover the block's warp count. // value per warp, thread 0 folds. kWarps must cover the block's warp count.
// With p.fold_ring, the last-finishing block additionally folds the final
// amax into the history window and publishes the next scale (atomicAdd
// ticket + fences), re-zeroing the amax slot and the counter for the next
// launch — the host-side delayed-scaling update chain disappears.
template <int kWarps> template <int kWarps>
__device__ __forceinline__ void publish_amax(float* amax, float v) { __device__ __forceinline__ void publish_amax(const FP8QuantizeParams& p,
float v) {
v = warp_reduce_max(v); v = warp_reduce_max(v);
__shared__ float slots[kWarps]; __shared__ float slots[kWarps];
const int tid = threadIdx.y * blockDim.x + threadIdx.x; const int tid = threadIdx.y * blockDim.x + threadIdx.x;
@@ -99,12 +123,23 @@ __device__ __forceinline__ void publish_amax(float* amax, float v) {
if (tid == 0) { if (tid == 0) {
#pragma unroll #pragma unroll
for (int w = 1; w < kWarps; ++w) v = fmaxf(v, slots[w]); for (int w = 1; w < kWarps; ++w) v = fmaxf(v, slots[w]);
atomic_max_float(amax, v); atomic_max_float(p.amax, v);
if (!p.fold_ring) return;
__threadfence();
const unsigned int ticket = atomicAdd(p.done, 1u);
__threadfence();
if (ticket != gridDim.x - 1u) return;
p.hist[p.hist_idx] = *p.amax;
float peak = p.hist[0];
for (int i = 1; i < p.hist_len; ++i) peak = fmaxf(peak, p.hist[i]);
*p.scale_out = fmaxf(peak / p.fp8_max / p.pow2_margin, 1e-12f);
*p.amax = 0.0f;
*p.done = 0u;
} }
} }
// Elementwise quantize kernel (out_layout 0): vectorized 16B loads -> fp8 // Elementwise quantize kernel (QuantLayout::RowMajor): vectorized 16B loads
// stores, fused amax over raw values. // -> fp8 stores, fused amax over raw values.
template <FP8Format Fmt, typename InT> template <FP8Format Fmt, typename InT>
__global__ void fp8_quantize_kernel(FP8QuantizeParams p) { __global__ void fp8_quantize_kernel(FP8QuantizeParams p) {
const float mult = *p.scale; const float mult = *p.scale;
@@ -155,82 +190,119 @@ __global__ void fp8_quantize_kernel(FP8QuantizeParams p) {
local_amax = fmaxf(local_amax, fabsf(v)); local_amax = fmaxf(local_amax, fabsf(v));
x8[i] = cvt_fp8<Fmt>(v * mult); x8[i] = cvt_fp8<Fmt>(v * mult);
} }
if (p.amax) publish_amax<8>(p.amax, local_amax); if (p.amax) publish_amax<8>(p, local_amax);
} }
// Tiled transpose quantize (out_layout 1/2): reads the [rows][cols] input // Tiled transpose quantize (QuantLayout::Transposed/Dual): reads the
// [rows][cols] input
// once and writes the fp8 bytes transposed ([cols][rows], so the contract // once and writes the fp8 bytes transposed ([cols][rows], so the contract
// dim lands K-contiguous for NT GEMM operands) and, in mode 2, the row-major // dim lands K-contiguous for NT GEMM operands) and, in mode 2, the row-major
// copy too. A 32x32 tile stages through shared memory: loads and writes // copy too. 64x32 tiles, one native pair load per row (a full 128B warp
// both stay coalesced, and the byte-wide staging is conflict-free — the +4 // read); rows whose pair is unaligned or ragged (odd widths, misaligned
// pad makes the store stride 9 words (coprime with the 32 banks) and the // bases) fall back to element loads in place. Staging goes through a byte
// read is a 32-byte broadcast segment. (A 64x64 split-half variant measured // tile whose pitch keeps the store stride coprime with the 32 banks.
// +21% L2-resident but -3..5% DRAM-bound; the real step mix ties, so the // (+25-35% over the former 32x32 scalar kernel on sub-4M tensors; ~5%
// simpler tile stays.) // slower once DRAM-saturated — accepted for the single-kernel shape.)
template <FP8Format Fmt, typename InT> template <FP8Format Fmt, typename InT>
__global__ void fp8_quantize_tiled_kernel(FP8QuantizeParams p) { __global__ void fp8_quantize_tiled_kernel(FP8QuantizeParams p) {
constexpr int kTile = 32; constexpr int kTileC = 64, kTileR = 32;
__shared__ uint8_t tile[kTile][kTile + 4]; // 34B pitch: staging stride is 17 words (coprime with the 32 banks) so
// the pair-byte stores stay conflict-free, and the byte-wise consume
// reads still span distinct words.
__shared__ uint8_t tile[kTileC][kTileR + 2];
const float mult = *p.scale; const float mult = *p.scale;
const auto* x = static_cast<const InT*>(p.input_ptr); const auto* x = static_cast<const InT*>(p.input_ptr);
const int r0 = blockIdx.y * kTile; const int r0 = blockIdx.y * kTileR;
const int c0 = blockIdx.x * kTile; const int c0 = blockIdx.x * kTileC;
const int r = r0 + threadIdx.y * 4; const int r = r0 + threadIdx.y * 4;
const int c = c0 + threadIdx.x; const int c = c0 + threadIdx.x * 2; // cols even => the pair is in-bounds
uint8_t q[4]; uint8_t q[4][2];
float local_amax = 0.0f; float local_amax = 0.0f;
// Vectorize the pair when both elements are in-bounds and the native
// 2-element load is aligned; odd widths, misaligned bases and ragged
// edges fall back to element loads row by row.
constexpr int kPairAlign = 2 * (int)sizeof(InT);
#pragma unroll #pragma unroll
for (int j = 0; j < 4; ++j) { for (int j = 0; j < 4; ++j) {
q[j] = 0; q[j][0] = 0;
q[j][1] = 0;
if (r + j < p.rows && c < p.cols) { if (r + j < p.rows && c < p.cols) {
const float v = const InT* a = x + (int64_t)(r + j) * p.cols + c;
quant_in_traits<InT>::to_float(x[(int64_t)(r + j) * p.cols + c]); if (c + 1 < p.cols &&
local_amax = fmaxf(local_amax, fabsf(v)); (reinterpret_cast<uintptr_t>(a) & (kPairAlign - 1)) == 0) {
q[j] = cvt_fp8<Fmt>(v * mult); float f[2];
quant_in_traits<InT>::load_pair(a, f);
#pragma unroll
for (int k = 0; k < 2; ++k) {
local_amax = fmaxf(local_amax, fabsf(f[k]));
q[j][k] = cvt_fp8<Fmt>(f[k] * mult);
}
} else {
const float v0 = quant_in_traits<InT>::to_float(a[0]);
local_amax = fmaxf(local_amax, fabsf(v0));
q[j][0] = cvt_fp8<Fmt>(v0 * mult);
if (c + 1 < p.cols) {
const float v1 = quant_in_traits<InT>::to_float(a[1]);
local_amax = fmaxf(local_amax, fabsf(v1));
q[j][1] = cvt_fp8<Fmt>(v1 * mult);
}
}
} }
} }
if (p.out_layout == 2) { if (p.out_layout == QuantLayout::Dual) {
uint8_t* out = static_cast<uint8_t*>(p.output_ptr); uint8_t* out = static_cast<uint8_t*>(p.output_ptr);
#pragma unroll #pragma unroll
for (int j = 0; j < 4; ++j) for (int j = 0; j < 4; ++j)
if (r + j < p.rows && c < p.cols) if (r + j < p.rows && c < p.cols) {
out[(int64_t)(r + j) * p.cols + c] = q[j]; uint8_t* o = out + (int64_t)(r + j) * p.cols + c;
const int64_t off = (int64_t)(r + j) * p.cols + c;
if (c + 1 < p.cols && (off & 1) == 0)
*reinterpret_cast<unsigned short*>(o) =
(unsigned short)(q[j][0] | (q[j][1] << 8));
else {
o[0] = q[j][0];
if (c + 1 < p.cols) o[1] = q[j][1];
}
}
} }
#pragma unroll #pragma unroll
for (int j = 0; j < 4; ++j) tile[threadIdx.x][threadIdx.y * 4 + j] = q[j]; for (int j = 0; j < 4; ++j)
#pragma unroll
for (int k = 0; k < 2; ++k)
tile[threadIdx.x * 2 + k][threadIdx.y * 4 + j] = q[j][k];
__syncthreads(); __syncthreads();
// Transposed scatter: output element (c, r) lives at c * rows + r; r // Transposed scatter: output element (c, r) lives at c * rows + r;
// tracks threadIdx.x so each warp writes one contiguous run. tile was // threadIdx.x tracks r so each warp writes one contiguous run. tile is
// written as tile[col][row], so input (r0+tx, c0+ty*4+j) reads back // [col][row]; warp y walks 8 columns, threads read down one column.
// from tile[ty*4+j][tx].
uint8_t* out_t = static_cast<uint8_t*>(p.output_transposed_ptr); uint8_t* out_t = static_cast<uint8_t*>(p.output_transposed_ptr);
#pragma unroll #pragma unroll
for (int j = 0; j < 4; ++j) { for (int i = 0; i < 8; ++i) {
const int oc = c0 + threadIdx.y * 4 + j; const int oc = c0 + threadIdx.y * 8 + i;
if (oc < p.cols && r0 + threadIdx.x < p.rows) if (oc < p.cols && r0 + threadIdx.x < p.rows)
out_t[(int64_t)oc * p.rows + r0 + threadIdx.x] = out_t[(int64_t)oc * p.rows + r0 + threadIdx.x] =
tile[threadIdx.y * 4 + j][threadIdx.x]; tile[threadIdx.y * 8 + i][threadIdx.x];
} }
if (p.amax) publish_amax<8>(p.amax, local_amax); if (p.amax) publish_amax<8>(p, local_amax);
} }
// Unified quantize launcher: Tiled selects the transpose kernel (out_layout // Unified quantize launcher: Tiled selects the transpose kernel
// 1/2) over the vectorized elementwise one. // (QuantLayout::Transposed/Dual) over the vectorized elementwise one. The
// transpose kernel vectorizes
// pair loads in-kernel and falls back to scalar loads at unaligned/ragged
// rows, so the host side picks only the grid.
template <FP8Format Fmt, typename InT, bool Tiled = false> template <FP8Format Fmt, typename InT, bool Tiled = false>
void launch_fp8_quantize(const FP8QuantizeParams& p, cudaStream_t stream) { void launch_fp8_quantize(const FP8QuantizeParams& p, cudaStream_t stream) {
if constexpr (Tiled) { if constexpr (Tiled) {
const dim3 grid((p.cols + 31) / 32, (p.rows + 31) / 32); const dim3 grid((p.cols + 63) / 64, (p.rows + 31) / 32);
if (grid.x == 0 || grid.y == 0) return; if (grid.x == 0 || grid.y == 0) return;
fp8_quantize_tiled_kernel<Fmt, InT> fp8_quantize_tiled_kernel<Fmt, InT><<<grid, dim3(32, 8), 0, stream>>>(p);
<<<grid, dim3(32, 8), 0, stream>>>(p);
} else { } else {
constexpr int kThreads = 256; constexpr int kThreads = 256;
constexpr int kVecElems = quant_in_traits<InT>::kVecElems; constexpr int kVecElems = quant_in_traits<InT>::kVecElems;
// One block per 256 vectors; at least one block so a tiny or // Grid-stride loops: any grid >= 1 is correct; one block per 256
// misaligned tensor's scalar tail is still covered. // vectors plus the tail block covers tiny and misaligned tensors.
int64_t blocks = (p.total / kVecElems + kThreads - 1) / kThreads; const int64_t blocks = 1 + p.total / (kVecElems * kThreads);
if (blocks < 1) blocks = 1;
fp8_quantize_kernel<Fmt, InT><<<blocks, kThreads, 0, stream>>>(p); fp8_quantize_kernel<Fmt, InT><<<blocks, kThreads, 0, stream>>>(p);
} }
} }
+41 -10
View File
@@ -4,7 +4,7 @@
- [Class Diagram](#class-diagram) — Full Mermaid class diagram across 10+ namespaces - [Class Diagram](#class-diagram) — Full Mermaid class diagram across 10+ namespaces
- [Module Overview](#module-overview) — Component inventory per module - [Module Overview](#module-overview) — Component inventory per module
- [Design Patterns](#design-patterns) — 15 documented patterns with classes - [Design Patterns](#design-patterns) — 16 documented patterns with classes
- [Core Relationships](#core-relationships) — 11 key inter-component relationships - [Core Relationships](#core-relationships) — 11 key inter-component relationships
## Class Diagram ## Class Diagram
@@ -816,8 +816,8 @@ classDiagram
class Executor { class Executor {
+AutoModel model +AutoModel model
+AutoTokenizer tokenizer
+PagePool kv_cache +PagePool kv_cache
+TaskCacheManager task_cache
+InferenceWorkspace _workspace +InferenceWorkspace _workspace
+Optional[str] device +Optional[str] device
+Optional[torch.dtype] dtype +Optional[torch.dtype] dtype
@@ -845,6 +845,7 @@ classDiagram
class InferenceScheduler { class InferenceScheduler {
+PagePool _cache +PagePool _cache
+TaskCacheManager _task_cache
+Executor _executor +Executor _executor
+TaskManager _task_mgr +TaskManager _task_mgr
+Event _stop_event +Event _stop_event
@@ -888,6 +889,24 @@ classDiagram
+release(pages) +release(pages)
} }
class AllocationStrategy {
<<abstract>>
+alloc(state, prompt_ids) bool
+free(state)
+extend(state, pos) bool
+write_indices(state, prompt_ids)
+record_hashes(state, prompt_ids, start_logical_page)
}
class ContiguousStrategy {
+write_indices(state, prompt_ids)
}
class PagedStrategy {
-Allocator _alloc
-RadixCache _prefix
}
class KVStorage { class KVStorage {
+int size +int size
+Tensor k_buffer +Tensor k_buffer
@@ -926,14 +945,21 @@ classDiagram
+bool contiguous +bool contiguous
-KVStorage _storage -KVStorage _storage
-ReqToTokenPool _req_pool -ReqToTokenPool _req_pool
-Allocator _alloc -AllocationStrategy _strategy
-RadixCache _prefix +strategy AllocationStrategy
+req_pool ReqToTokenPool
+bind_tasks(req_indices, seq_lens, workspace, device, start_pos, incremental) KVCache
}
class TaskCacheManager {
-PagePool _pool
-Dict _states
+task_alloc(task_id, prompt_ids) bool +task_alloc(task_id, prompt_ids) bool
+task_free(task_id) +task_free(task_id)
+task_extend(task_id, pos) bool +task_extend(task_id, pos) bool
+task_cached(task_id) int +task_cached(task_id) int
+task_record_hashes(task_id, prompt_ids, start_logical_page) +task_record_hashes(task_id, prompt_ids, start_logical_page)
+bind_tasks(task_ids, workspace, device, start_pos) KVCache +bind(task_ids, workspace) KVCache
} }
class Task { class Task {
@@ -1316,17 +1342,22 @@ classDiagram
PositionIdStrategy <|-- DocResetPositionId PositionIdStrategy <|-- DocResetPositionId
PositionIdStrategy <|-- ContinuousPositionId PositionIdStrategy <|-- ContinuousPositionId
StoreWriter <|-- BinWriter StoreWriter <|-- BinWriter
AllocationStrategy <|-- ContiguousStrategy
AllocationStrategy <|-- PagedStrategy
RawRollout <|-- RolloutResult RawRollout <|-- RolloutResult
LaunchStrategy <|-- TorchrunStrategy LaunchStrategy <|-- TorchrunStrategy
LaunchStrategy <|-- LocalStrategy LaunchStrategy <|-- LocalStrategy
%% --- Composition (strong ownership, part destroyed with whole) --- %% --- Composition (strong ownership, part destroyed with whole) ---
PagePool *-- KVStorage PagePool *-- KVStorage
PagePool *-- ReqToTokenPool PagePool *-- ReqToTokenPool
PagePool *-- Allocator PagePool *-- AllocationStrategy
PagePool *-- RadixCache PagedStrategy *-- Allocator
PagedStrategy *-- RadixCache
TaskCacheManager o-- PagePool
RadixCache *-- RadixNode RadixCache *-- RadixNode
InferenceEngine *-- InferenceScheduler InferenceEngine *-- InferenceScheduler
InferenceScheduler *-- PagePool InferenceScheduler *-- PagePool
InferenceScheduler *-- TaskCacheManager
InferenceScheduler *-- Executor InferenceScheduler *-- Executor
Executor *-- InferenceWorkspace Executor *-- InferenceWorkspace
InferenceScheduler *-- TaskManager InferenceScheduler *-- TaskManager
@@ -1419,7 +1450,7 @@ classDiagram
Task --> TaskStatus Task --> TaskStatus
InferenceEngine --> AutoModel InferenceEngine --> AutoModel
Executor --> AutoModel Executor --> AutoModel
Executor --> AutoTokenizer Executor --> TaskCacheManager
TaskManager --> AutoTokenizer TaskManager --> AutoTokenizer
``` ```
@@ -1436,7 +1467,7 @@ classDiagram
| **astrai.model** | ModelFactory, AutoModel, AutoRegressiveLM, EmbeddingEncoder, DecoderBlock, GQA, MLA, MLP, DeepSeekMoE, AttnFactory, FFNFactory, RMSNorm, Linear, LoRAConfig, LoRALinear, RotaryEmbedding, Embedding | Neural network model | | **astrai.model** | ModelFactory, AutoModel, AutoRegressiveLM, EmbeddingEncoder, DecoderBlock, GQA, MLA, MLP, DeepSeekMoE, AttnFactory, FFNFactory, RMSNorm, Linear, LoRAConfig, LoRALinear, RotaryEmbedding, Embedding | Neural network model |
| **astrai.tokenize** | AutoTokenizer, ChatTemplate | Tokenizer and chat template | | **astrai.tokenize** | AutoTokenizer, ChatTemplate | Tokenizer and chat template |
| **astrai.trainer** | Trainer, TrainContext, TrainContextBuilder, BaseStrategyGRPOStrategy, StrategyFactory, BaseSchedulerWSDScheduler, SchedulerFactory, TrainCallback(Protocol)MetricCallback, CallbackFactory, RawRollout, RolloutResult, BaseRewardModel, RolloutGenerator, RolloutRunner | Training workflow | | **astrai.trainer** | Trainer, TrainContext, TrainContextBuilder, BaseStrategyGRPOStrategy, StrategyFactory, BaseSchedulerWSDScheduler, SchedulerFactory, TrainCallback(Protocol)MetricCallback, CallbackFactory, RawRollout, RolloutResult, BaseRewardModel, RolloutGenerator, RolloutRunner | Training workflow |
| **astrai.inference** | InferenceEngine, InferenceScheduler, Executor, InferenceWorkspace, PagePool, KVStorage, ReqToTokenPool, KVCache, Allocator, RadixCache, Task, TaskManager, TaskStatus, StreamDecoder, GenerateResult, BaseSamplingStrategySamplingPipeline, FrequencyPenaltyStrategy, ProtocolHandler, ResponseBuilder, OpenAIResponseBuilder, AnthropicResponseBuilder, StopChecker, GenContext, StopInfo, ChatMessage, FunctionDef, ToolDef, ChatCompletionRequest, AnthropicMessage, MessagesRequest, BaseToolParser, ToolParserFactory, SimpleJsonToolParser | Inference service | | **astrai.inference** | InferenceEngine, InferenceScheduler, Executor, InferenceWorkspace, PagePool, TaskCacheManager, KVStorage, ReqToTokenPool, KVCache, Allocator, RadixCache, AllocationStrategy, ContiguousStrategy, PagedStrategy, Task, TaskManager, TaskStatus, StreamDecoder, GenerateResult, BaseSamplingStrategySamplingPipeline, FrequencyPenaltyStrategy, ProtocolHandler, ResponseBuilder, OpenAIResponseBuilder, AnthropicResponseBuilder, StopChecker, GenContext, StopInfo, ChatMessage, FunctionDef, ToolDef, ChatCompletionRequest, AnthropicMessage, MessagesRequest, BaseToolParser, ToolParserFactory, SimpleJsonToolParser | Inference service |
| **astrai.extension** | `backend` policy package, `ops` kernel-wrapper package, `fp8.py` FP8 strategy layer, AttentionBackend, TorchNativeBackend, CudaBackend, FlashAttnBackend, attention, attn_backend, ATTN_BACKEND, apply_rotary_emb, is_available | Stable API over attention/rotary/FP8 execution policy and optional CUDA kernels | | **astrai.extension** | `backend` policy package, `ops` kernel-wrapper package, `fp8.py` FP8 strategy layer, AttentionBackend, TorchNativeBackend, CudaBackend, FlashAttnBackend, attention, attn_backend, ATTN_BACKEND, apply_rotary_emb, is_available | Stable API over attention/rotary/FP8 execution policy and optional CUDA kernels |
| **astrai.optim** | OptimizerFactory, MuonAdamW, NoraNadamW, ManoAdamW, composite_step/composite_zero_grad/composite_state_dict, partition_optimizer_parameters | Built-in optimizers (`muon_adamw` / `nora_nadamw` / `mano_adamw`) with shared composite-optimizer helpers | | **astrai.optim** | OptimizerFactory, MuonAdamW, NoraNadamW, ManoAdamW, composite_step/composite_zero_grad/composite_state_dict, partition_optimizer_parameters | Built-in optimizers (`muon_adamw` / `nora_nadamw` / `mano_adamw`) with shared composite-optimizer helpers |
| **astrai.parallel** | spawn_parallel_fn, setup_parallel, get_rank/get_world_size/get_current_device, only_on_rank, LaunchStrategy, TorchrunStrategy, LocalStrategy, BaseExecutor, ExecutorFactory, NoneExecutor, DDPExecutor, FSDPExecutor, GradientState, AccumOptimizer, AccumScheduler | Distributed parallel & gradient accumulation | | **astrai.parallel** | spawn_parallel_fn, setup_parallel, get_rank/get_world_size/get_current_device, only_on_rank, LaunchStrategy, TorchrunStrategy, LocalStrategy, BaseExecutor, ExecutorFactory, NoneExecutor, DDPExecutor, FSDPExecutor, GradientState, AccumOptimizer, AccumScheduler | Distributed parallel & gradient accumulation |
@@ -1478,4 +1509,4 @@ classDiagram
10. **AutoModel**: `from_pretrained()` loads `config.json` + `model.safetensors`, `_disable_random_init` replaces `nn.init.*` with no-ops 10. **AutoModel**: `from_pretrained()` loads `config.json` + `model.safetensors`, `_disable_random_init` replaces `nn.init.*` with no-ops
11. **Protocols**: `OptimizerProtocol` / `SchedulerProtocol` — structural subtyping for `AccumOptimizer` / `AccumScheduler` wrappers 11. **Protocols**: `OptimizerProtocol` / `SchedulerProtocol` — structural subtyping for `AccumOptimizer` / `AccumScheduler` wrappers
> Document Update Time: 2026-08-22 > Document Update Time: 2026-08-29
+3 -3
View File
@@ -305,7 +305,7 @@ cycle belong under `TYPE_CHECKING`.
- **`AttentionBackend`** (ABC): `fwd_decode` / `fwd_prefill` abstract methods, `forward` dispatches by q_len - **`AttentionBackend`** (ABC): `fwd_decode` / `fwd_prefill` abstract methods, `forward` dispatches by q_len
- **`CudaBackend`**: CUDA kernel dispatch — decode via `attn_paged_decode` (page_size=1), prefill via `attn_paged_prefill` (ragged batch, `qo_indptr` + `kv_indptr`). Default on GPU. - **`CudaBackend`**: CUDA kernel dispatch — decode via `attn_paged_decode` (page_size=1), prefill via `attn_paged_prefill` (ragged batch, `qo_indptr` + `kv_indptr`). Default on GPU.
- **`FlashAttnBackend`**: Optional flash-attn dispatch with `flash_attn_with_kvcache` fast path. - **`FlashAttnBackend`**: Optional flash-attn dispatch via `flash_attn_varlen_func` over gathered flat K/V.
- **`TorchNativeBackend`**: SDPA with indirect KV cache gather (always-available fallback) - **`TorchNativeBackend`**: SDPA with indirect KV cache gather (always-available fallback)
Default priority: cuda > flash > torch. Set ``ASTR_BACKEND=cuda|torch_native|flash`` Default priority: cuda > flash > torch. Set ``ASTR_BACKEND=cuda|torch_native|flash``
@@ -415,7 +415,7 @@ nvcc -I csrc -arch=sm_89 -O3 --use_fast_math \
Test files: Test files:
- `attn_test.cu` — decode + prefill kernels (correctness tables + benchmarks) - `attn_test.cu` — decode + prefill kernels (correctness tables + benchmarks)
- `attn_paged_test.cu` — paged decode/prefill kernels - `attn_paged_test.cu` — paged decode/prefill kernels
- `fp8_mma_test.cu` — BF16→FP8→BF16 MMA demo (sm_89) - `fp8_test.cu` — single-warp bf16→fp8→mma.sync sanity check + full FP8 GEMM correctness (sm_89)
## Benchmarks ## Benchmarks
@@ -482,4 +482,4 @@ csrc/
Compiled `.so` files are placed in `astrai/extension/lib/`, separate from Python source files. Compiled `.so` files are placed in `astrai/extension/lib/`, separate from Python source files.
> Document Update Time: 2026-08-22 > Document Update Time: 2026-08-29
+83 -6
View File
@@ -42,10 +42,11 @@ runtime:
stop_timeout_seconds: 600 stop_timeout_seconds: 600
checkpoint_keep_last: 5 checkpoint_keep_last: 5
# max_duration_hours: 12 # max_duration_hours: 12
# Add host-specific workarounds only when required: # Optional; entries are passed verbatim into the trainer container
# (see "Per-Job Environment"):
# environment: # environment:
# NCCL_P2P_DISABLE: "1" # ASTR_LOG_LEVEL: DEBUG
# NCCL_NET_GDR_LEVEL: "0" # ASTR_BACKEND: torch_native
``` ```
- Relative paths resolve from the YAML file's directory, not the current shell. - Relative paths resolve from the YAML file's directory, not the current shell.
@@ -57,11 +58,85 @@ runtime:
Use `fsdp` explicitly when model sharding is required. Use `fsdp` explicitly when model sharding is required.
- To select specific physical GPUs, replace `all` with a list such as - To select specific physical GPUs, replace `all` with a list such as
`devices: [0, 1]`. `devices: [0, 1]`.
- `environment` values are explicitly passed to the training container. Keep - `environment` entries apply only to the job defined by this YAML file, not to
host-specific NCCL workarounds here; they are not universal defaults. the host or to other jobs. Keep the section omitted unless this job's GPU
selection needs it; see [Per-Job Environment](#per-job-environment).
- `max_duration_hours` starts a detached host timer that calls the same graceful - `max_duration_hours` starts a detached host timer that calls the same graceful
`stop` command. A manual stop cancels the timer. `stop` command. A manual stop cancels the timer.
## Per-Job Environment
`runtime.environment` is scoped to one job. `start` passes only the entries of
the config file it was given, so a variable reaches exactly the GPUs declared
in that file's `runtime.gpu.devices` and nothing else. Two jobs on the same
machine can therefore differ: a job whose GPUs have working peer-to-peer keeps
the section omitted, a job whose GPUs cross broken PCIe/NVLink paths declares
the NCCL workarounds, and a job on an NVSwitch fabric can pin the NVLink fast
path on.
Because of that scoping, the effective pattern is one YAML per GPU group
rather than one shared YAML that gets edited whenever the device list changes:
```yaml
# train-local.yaml: GPUs with working peer-to-peer; nothing to declare
runtime:
gpu:
devices: [0, 1]
# train-cross-pcie.yaml: this GPU set crosses broken paths, so only this job
# declares the workarounds (confirm first; see docs/guides/distributed.md)
runtime:
gpu:
devices: [4, 5, 6, 7]
environment:
NCCL_P2P_DISABLE: "1"
NCCL_NET_GDR_LEVEL: "0"
```
The same mechanism carries positive tuning, not just workarounds. On an
NVSwitch node (Hopper-class GPUs with fabric manager running), NVLink SHARP
multicast (NVLS) is the fast allreduce path and NCCL enables it automatically
where supported. A job may pin it on explicitly and raise channel parallelism
when benchmarks show the NVLink bandwidth is underused:
```yaml
# train-nvlink.yaml: NVSwitch node; keep the disables OUT and pin the fast
# path on instead (verify support with NCCL_DEBUG=INFO first)
runtime:
gpu:
devices: [0, 1, 2, 3]
environment:
NCCL_NVLS_ENABLE: "1"
NCCL_MIN_NCHANNELS: "8"
# NCCL_ALGO: NVLS # force one algorithm; unsupported values fail loudly
```
NVLS requires NVSwitch multicast support; on plain NVLink bridges or PCIe-only
sets, keep the section omitted and let NCCL pick Ring/Tree with P2P. Newer
drivers list the actual interconnect and NVLS support directly in
`nvidia-smi topo -m`, so check that before assuming.
Confirm a variable is needed before adding it, and only in the YAML of the job
that hits the problem:
```bash
nvidia-smi topo -m # check P2P support between exactly the selected GPUs
NCCL_DEBUG=INFO # confirm NCCL transport errors before disabling them
```
See `docs/guides/distributed.md` for what each troubleshooting variable
disables. The two directions are mutually exclusive: `NCCL_P2P_DISABLE` and
`NCCL_NET_GDR_LEVEL` remove bandwidth and must never appear in the same
environment as the NVLink entries above.
Semantics:
- Values must be scalars and are rendered with `str()`, so quote them
explicitly (`"1"`, `"0"`) instead of relying on YAML booleans or numbers.
- A `null` value exports the name with an empty value.
- This section is the only path for extra host variables into the trainer
container; variables exported in the host shell do not pass through Compose.
## Fixed Container Paths ## Fixed Container Paths
| Runtime path | Container path | Access | | Runtime path | Container path | Access |
@@ -123,5 +198,7 @@ the Docker timeout expires.
3. Do not force DDP for a model that requires FSDP; declare the mode explicitly. 3. Do not force DDP for a model that requires FSDP; declare the mode explicitly.
4. Do not use `kill -9` for routine shutdown; use `scripts/train.sh stop CONFIG`. 4. Do not use `kill -9` for routine shutdown; use `scripts/train.sh stop CONFIG`.
5. The image user is built with the host UID/GID so mounted checkpoints retain usable ownership. 5. The image user is built with the host UID/GID so mounted checkpoints retain usable ownership.
6. Scope `runtime.environment` to the job YAML that needs it; do not copy NCCL
workarounds into every config.
> Document Update Time: 2026-08-22 > Document Update Time: 2026-08-29
+1 -1
View File
@@ -185,7 +185,7 @@ The extension package separates mechanism from policy:
Attention computation is decoupled from the model via `AttentionBackend` ABC (`astrai/extension/backend/attention.py`): Attention computation is decoupled from the model via `AttentionBackend` ABC (`astrai/extension/backend/attention.py`):
- **`CudaBackend`** (default when supported): decode path uses `attn_paged_decode` with `page_size=1` (the `req_to_token` table serves as the page table, each token slot is a single-token "page"); prefill path uses the ragged-batch `attn_paged_prefill` (addresses each request via `qo_indptr` + `kv_indptr` directly against the flat pool). - **`CudaBackend`** (default when supported): decode path uses `attn_paged_decode` with `page_size=1` (the `req_to_token` table serves as the page table, each token slot is a single-token "page"); prefill path uses the ragged-batch `attn_paged_prefill` (addresses each request via `qo_indptr` + `kv_indptr` directly against the flat pool).
- **`FlashAttnBackend`**: optional flash-attn dispatch with `flash_attn_with_kvcache` fast path for contiguous cache; falls back to KV gather + `flash_attn_func`. - **`FlashAttnBackend`**: optional flash-attn dispatch; inference paths gather flat K/V from the pool via `req_to_token` and call `flash_attn_varlen_func` over the ragged batch (fp16/bf16 only); dense mask-free training calls use `flash_attn_func`.
- **`TorchNativeBackend`** (always-available fallback): writes K/V to cache, gathers via `req_to_token` indirect indexing, calls `F.scaled_dot_product_attention`. - **`TorchNativeBackend`** (always-available fallback): writes K/V to cache, gathers via `req_to_token` indirect indexing, calls `F.scaled_dot_product_attention`.
- The `attention(...)` entry point uses cuda > flash > torch priority and chooses another compatible backend when an automatically selected backend cannot handle a call. - The `attention(...)` entry point uses cuda > flash > torch priority and chooses another compatible backend when an automatically selected backend cannot handle a call.
- Resolution precedence is: explicit `attn_backend(...)` context > `ASTR_BACKEND` env > default. An explicit `attn_backend(...)` selection is strict (incompatible calls raise); `ASTR_BACKEND` is a default-level override that falls back to a compatible backend when incapable. Training calls (`fwd=None`, no KV cache) resolve by capability: the CUDA cache kernels cannot run without a cache, so they fall back to flash (mask-free/causal calls only) and finally to torch SDPA. - Resolution precedence is: explicit `attn_backend(...)` context > `ASTR_BACKEND` env > default. An explicit `attn_backend(...)` selection is strict (incompatible calls raise); `ASTR_BACKEND` is a default-level override that falls back to a compatible backend when incapable. Training calls (`fwd=None`, no KV cache) resolve by capability: the CUDA cache kernels cannot run without a cache, so they fall back to flash (mask-free/causal calls only) and finally to torch SDPA.
+3 -2
View File
@@ -190,8 +190,9 @@ python scripts/tools/train.py \
```bash ```bash
export CUDA_VISIBLE_DEVICES=0,1,2,3 export CUDA_VISIBLE_DEVICES=0,1,2,3
export NCCL_P2P_DISABLE=1 # Only if this host's NCCL transport is broken; see docs/guides/distributed.md:
export NCCL_NET_GDR_LEVEL=0 # export NCCL_P2P_DISABLE=1
# export NCCL_NET_GDR_LEVEL=0
python scripts/tools/train.py \ python scripts/tools/train.py \
--train_type=seq \ --train_type=seq \
+89 -44
View File
@@ -2,8 +2,9 @@
The kernel-level tests exercise the two stateless primitives (``quantize`` for The kernel-level tests exercise the two stateless primitives (``quantize`` for
bf16/fp16/fp32 -> FP8, ``mm_fp8`` for the pre-quantized GEMM with transposed bf16/fp16/fp32 -> FP8, ``mm_fp8`` for the pre-quantized GEMM with transposed
operands); the policy-level tests (recipes, autocast context, per-tensor meta, operands); the policy-level tests (recipes, autocast context, per-tensor
CPU fallbacks of the custom ops) run without a GPU. meta) run without a GPU. The primitives themselves are CUDA-only
(attention-style direct wrappers no torch.library dispatch layer).
""" """
import threading import threading
@@ -14,9 +15,8 @@ import torch.nn.functional as F
import astrai.extension.fp8 as f8mod import astrai.extension.fp8 as f8mod
from astrai.extension.fp8 import ( from astrai.extension.fp8 import (
DelayedScaling,
DynamicScaling,
FP8Format, FP8Format,
FP8Recipe,
FP8TensorMeta, FP8TensorMeta,
_ScaleRing, _ScaleRing,
fp8_autocast, fp8_autocast,
@@ -24,7 +24,7 @@ from astrai.extension.fp8 import (
fp8_linear_enabled, fp8_linear_enabled,
fp8_state, fp8_state,
) )
from astrai.extension.ops.fp8 import mm_fp8, quantize from astrai.extension.ops.fp8 import mm_fp8, quantize, quantize_dual
from tests.conftest import skip_no_fp8 from tests.conftest import skip_no_fp8
@@ -233,7 +233,7 @@ def test_delayed_scaling_forward_uses_snapshot_scale():
dev = torch.device("cuda") dev = torch.device("cuda")
state = f8mod.fp8_state() state = f8mod.fp8_state()
state.reset() state.reset()
state.default_recipe = DelayedScaling(history_len=1, margin=0) state.default_recipe = FP8Recipe(history_len=1, margin=0)
state.default_format = FP8Format.E4M3 state.default_format = FP8Format.E4M3
try: try:
m, n, k = 32, 16, 64 m, n, k = 32, 16, 64
@@ -272,7 +272,7 @@ def test_fp8_linear_forward_and_backward():
state = f8mod.fp8_state() state = f8mod.fp8_state()
state.reset() state.reset()
state.default_recipe = DynamicScaling() state.default_recipe = FP8Recipe(dynamic=True)
try: try:
out, _, _ = f8mod.fp8_linear_forward(x, weight, bias) out, _, _ = f8mod.fp8_linear_forward(x, weight, bias)
@@ -399,13 +399,13 @@ def test_mm_fp8_matches_scaled_mm():
def test_recipe_scale_from_history(): def test_recipe_scale_from_history():
"""Delayed: max over the window + margin; dynamic: current amax.""" """Delayed: max over the window + margin; dynamic: current amax."""
hist = torch.tensor([1.0, 2.0, 0.5]) hist = torch.tensor([1.0, 2.0, 0.5])
d = DelayedScaling(history_len=3, margin=0) d = FP8Recipe(history_len=3, margin=0)
assert torch.allclose(d.scale_from_history(hist, "e4m3"), torch.tensor(2.0 / 448.0)) assert torch.allclose(d.scale_from_history(hist, "e4m3"), torch.tensor(2.0 / 448.0))
d_m = DelayedScaling(history_len=3, margin=2) d_m = FP8Recipe(history_len=3, margin=2)
assert torch.allclose( assert torch.allclose(
d_m.scale_from_history(hist, "e4m3"), torch.tensor(2.0 / 448.0 / 4.0) d_m.scale_from_history(hist, "e4m3"), torch.tensor(2.0 / 448.0 / 4.0)
) )
dyn = DynamicScaling() dyn = FP8Recipe(dynamic=True)
amax = torch.tensor([0.25]) amax = torch.tensor([0.25])
assert torch.allclose( assert torch.allclose(
dyn.scale_from_history(amax, "e4m3"), torch.tensor(0.25 / 448.0) dyn.scale_from_history(amax, "e4m3"), torch.tensor(0.25 / 448.0)
@@ -423,62 +423,107 @@ def test_fp8_format_enum():
def test_fp8_autocast_context(): def test_fp8_autocast_context():
"""fp8_autocast sets and restores recipe + format on the global state.""" """fp8_autocast pushes and restores the thread-local active config."""
state = fp8_state() state = fp8_state()
prev = (state.enabled, state.recipe, state.fp8_format) state.reset()
try: try:
with fp8_autocast(enabled=True, fp8_format="hybrid", update_interval=8): with fp8_autocast(enabled=True, fp8_format="hybrid", update_interval=8):
assert state.enabled cfg = f8mod._active_config.get()
assert isinstance(state.recipe, DelayedScaling) assert cfg is not None and cfg.enabled
assert state.recipe.history_len == 8 assert not cfg.recipe.dynamic
assert state.fp8_format is FP8Format.HYBRID assert cfg.recipe.history_len == 8
with fp8_autocast(enabled=True, recipe=DynamicScaling(), fp8_format="e4m3"): assert cfg.fp8_format is FP8Format.HYBRID
assert isinstance(state.recipe, DynamicScaling) with fp8_autocast(
assert state.fp8_format is FP8Format.E4M3 enabled=True, recipe=FP8Recipe(dynamic=True), fp8_format="e4m3"
assert state.fp8_format is FP8Format.HYBRID # restored on exit ):
assert not state.enabled inner = f8mod._active_config.get()
assert inner.recipe.dynamic
assert inner.fp8_format is FP8Format.E4M3
assert f8mod._active_config.get() is cfg # restored on exit
assert f8mod._active_config.get() is None
assert not fp8_linear_enabled()
finally: finally:
state.enabled, state.recipe, state.fp8_format = prev state.reset()
def test_fp8_tensor_meta_delayed_update(): def test_fp8_tensor_meta_delayed_update():
"""Meta seeds from data; hist/scale are packed views of one state buffer.""" """Meta seeds from data; hist/scale are packed views of one state buffer."""
meta = FP8TensorMeta(torch.device("cpu"), DelayedScaling(history_len=4, margin=0)) recipe = FP8Recipe(history_len=4, margin=0)
meta = FP8TensorMeta(
_ScaleRing(torch.device("cpu"), recipe),
_ScaleRing(torch.device("cpu"), recipe),
_ScaleRing(torch.device("cpu"), recipe),
)
w = torch.randn(8, 8) w = torch.randn(8, 8)
meta.w.seed(w, "e4m3") meta.w.seed(w, "e4m3")
assert meta.w.initialized assert meta.w.initialized
torch.testing.assert_close(meta.w.scale, (w.abs().amax() / 448.0).reshape(1)) torch.testing.assert_close(meta.w.scale, (w.abs().amax() / 448.0).reshape(1))
# [hist | scale] packing: views alias the single state buffer. # [hist | scale | legacy | amax | done] packing: views alias one buffer.
assert meta.w.state.numel() == 4 + 2 assert meta.w.state.numel() == 4 + 4
assert meta.w.hist.data_ptr() == meta.w.state.data_ptr() assert meta.w.hist.data_ptr() == meta.w.state.data_ptr()
assert meta.w.scale.data_ptr() == meta.w.state[4:].data_ptr() assert meta.w.scale.data_ptr() == meta.w.state[4:].data_ptr()
meta.w.advance() meta.w.advance()
assert meta.w.idx == 1 assert meta.w.idx == 1
# update folds a fresh amax into the window and publishes the next scale # fold_args hands the kernel the buffer, the slot and the recipe constants
amax = torch.tensor([8.0]) args = meta.w.fold_args("e4m3")
meta.w.update(amax, "e4m3") assert args["ring_state"] is meta.w.state and args["hist_idx"] == 1
torch.testing.assert_close(meta.w.scale, torch.tensor([8.0 / 448.0])) assert args["fp8_max"] == 448.0 and args["pow2_margin"] == 1.0
def test_quantize_cpu_fallback(): @skip_no_fp8
"""CPU fallback of the quantize primitive (scale semantics + amax).""" @pytest.mark.parametrize("fmt", ["e4m3", "e5m2"])
x = torch.randn(16, 32, dtype=torch.bfloat16) def test_quantize_dual_and_transposed_orientations(fmt):
scale = torch.tensor([0.5]) # quantize multiplier """quantize_dual yields both orientations from one read; quantize's
x8, amax = quantize(x, scale, "e4m3") transposed switch keeps the 2-tuple arity with the [cols][rows] layout."""
assert x8.dtype == torch.float8_e4m3fn torch.manual_seed(11)
ref = (x.float() * 0.5).to(torch.float8_e4m3fn) x = torch.randn(37, 67, device="cuda", dtype=torch.bfloat16) * 3
assert torch.equal(x8, ref) mult = _scale(x).reciprocal()
x8, amax = quantize(x, mult, fmt)
x8T, _ = quantize(x, mult, fmt, transposed=True)
d8, d8T, _ = quantize_dual(x, mult, fmt)
assert x8T.shape == (67, 37)
assert torch.equal(x8.view(torch.uint8), d8.view(torch.uint8))
assert torch.equal(x8T.view(torch.uint8), d8T.view(torch.uint8))
assert torch.equal(x8T.t().contiguous().view(torch.uint8), x8.view(torch.uint8))
torch.testing.assert_close(amax, x.abs().amax().float().reshape(1)) torch.testing.assert_close(amax, x.abs().amax().float().reshape(1))
def test_mm_fp8_cpu_fallback(): @skip_no_fp8
a8 = torch.tensor([[1.0, 2.0]], dtype=torch.float8_e4m3fn) @pytest.mark.parametrize("fmt,fmax", [("e4m3", 448.0), ("e5m2", 57344.0)])
b8 = torch.tensor([[3.0], [4.0]], dtype=torch.float8_e4m3fn) @pytest.mark.parametrize("margin", [0, 1])
scale = torch.tensor([1.0]) def test_quantize_ring_fold_matches_host_update(fmt, fmax, margin):
out = mm_fp8(a8, b8, scale) """The in-kernel delayed-scaling fold matches a host-side reference."""
ref = (a8.float() @ b8.float() * 1.0).to(torch.bfloat16) dev = torch.device("cuda")
torch.testing.assert_close(out, ref) n, idx = 4, 2
torch.manual_seed(3)
x = torch.randn(128, 96, dtype=torch.bfloat16, device=dev) * 3
mult = torch.tensor([0.01], device=dev)
pow2m = float(2**margin)
# Reference: legacy quantize + the host fold it used to return amax for.
x8_ref, amax = quantize(x, mult, fmt)
hist = torch.full((n,), 1.0, device=dev)
hist[idx] = amax.to(torch.float32)
scale = (hist.max() / fmax / pow2m).clamp_min(1e-12).reshape(1)
# Fused: same window, fold inside the quantize kernel's last block.
ring = torch.zeros(n + 4, device=dev)
ring[:n].fill_(1.0)
x8, _ = quantize(
x,
mult,
fmt,
ring_state=ring,
hist_idx=idx,
fp8_max=fmax,
pow2_margin=pow2m,
)
assert torch.equal(x8.view(torch.uint8), x8_ref.view(torch.uint8))
torch.testing.assert_close(ring[:n], hist, rtol=0, atol=0)
torch.testing.assert_close(ring[n : n + 1], scale, rtol=0, atol=0)
assert float(ring[n + 2]) == 0.0 # amax slot self-cleaned
assert int(ring[n + 3].view(torch.int32)) == 0 # done counter reset
# -------------------------------------------------------------------------- # --------------------------------------------------------------------------
+46 -1
View File
@@ -393,7 +393,9 @@ def test_decode_does_not_reuse_previous_batch_state():
old_info = object() old_info = object()
new_info = object() new_info = object()
executor._decode_cache = DecodeSteadyState(("old",), [2], old_info) executor._decode_cache = DecodeSteadyState(("old",), [2], old_info)
executor._sample_logits = MagicMock(return_value=[3]) executor._sample_logits = MagicMock(
return_value=([3], torch.tensor([3], dtype=torch.long))
)
task = Task("new", list(range(8)), temperature=0) task = Task("new", list(range(8)), temperature=0)
task.input_tokens = 8 task.input_tokens = 8
@@ -412,3 +414,46 @@ def test_decode_does_not_reuse_previous_batch_state():
args, kwargs = executor._sample_logits.call_args args, kwargs = executor._sample_logits.call_args
assert args[1:] == ([task], False) assert args[1:] == ([task], False)
assert kwargs["info"] is new_info assert kwargs["info"] is new_info
def test_decode_fills_input_ids_from_device_on_matching_signature():
"""Steady-state decode copies cached device tokens, skipping the host."""
executor = object.__new__(Executor)
executor.device = torch.device("cpu")
executor.task_cache = MagicMock()
executor.task_cache.bind_was_steady = True
executor.task_cache.bind.return_value = MagicMock()
executor._graph_supported = False
executor._graph_ctx = SimpleNamespace(enabled=False)
workspace = MagicMock()
workspace.position_ids = torch.tensor([2], dtype=torch.long)
workspace.fill_input_ids_from_device.return_value = torch.tensor(
[9], dtype=torch.long
)
executor._workspace = workspace
executor.model = MagicMock(
return_value={"logits": torch.zeros(1, 1, 10, dtype=torch.float32)}
)
info = object()
tokens = torch.tensor([3], dtype=torch.long)
executor._decode_cache = DecodeSteadyState(("t1",), [2], info, last_tokens=tokens)
executor._sample_logits = MagicMock(return_value=([3], tokens))
task = Task("t1", list(range(8)), temperature=0)
task.input_tokens = 8
task.output_ids = [7]
task.mark_prefill_done()
with patch(
"astrai.inference.runtime.executor._build_sampling_batch_info",
return_value=info,
):
assert executor.execute_decode([task]) == [3]
workspace.fill_input_ids.assert_not_called()
workspace.fill_input_ids_from_device.assert_called_once_with(tokens)
assert workspace.position_ids.tolist() == [3]
assert executor._decode_cache.task_sig == ("t1",)
assert executor._decode_cache.last_tokens is tokens