From 4dc5e923e06f38b3f19c2827c4c075cb9ad40f9b Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Tue, 25 Aug 2026 11:07:31 +0800 Subject: [PATCH] perf: finalize fp8 scale rings inside quantize kernels MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - last-block epilogue (threadfence + counter elect) folds amax into hist[idx], reduces the window and publishes the next scale on device — zero extra launches per linear layer - _ScaleRing packs [hist | scale | counter] into one CUDA buffer; the eager hist-write / max / scale-copy chain and update() are gone - split FP8QuantizeParams out of FP8Params so each operator owns its fields; linear_forward/backward_fp8 take optional ring arguments - e2e 12L/dim1024/B4xT512 (fused AdamW): fp8 137.8ms/step vs bf16 210.3ms, 1.53x; fwd 1.82x, bwd 1.50x --- astrai/extension/fp8.py | 83 +++++++++++++++++++++--------- astrai/extension/ops/fp8.py | 62 ++++++++++++++++++++-- csrc/kernels/fp8/common.h | 62 +++++++++++++++------- csrc/kernels/fp8/gemm.cuh | 49 +++++++++++++++++- csrc/kernels/fp8/ops.cu | 91 ++++++++++++++++++++++----------- tests/extension/test_fp8_mma.py | 61 ++++++++++++++++++++-- 6 files changed, 324 insertions(+), 84 deletions(-) diff --git a/astrai/extension/fp8.py b/astrai/extension/fp8.py index 76c7752..715477a 100644 --- a/astrai/extension/fp8.py +++ b/astrai/extension/fp8.py @@ -105,26 +105,32 @@ class DynamicScaling(FP8Recipe): class _ScaleRing: - """One operand's delayed-scaling state: amax history ring + derived scale. + """One operand's delayed-scaling state, packed for in-kernel finalization. - The ring captures its recipe at construction; ``update`` records a fresh - amax and refreshes the scale for the *next* step (delayed one step). + ``state`` is a single float32 CUDA buffer ``[hist[n] | scale | counter]`` + (``hist`` / ``scale`` are views). The quantize kernel's last-finishing + block records the freshly measured amax into ``hist[idx]``, reduces the + window and publishes the next step's scale entirely on device — the + Python-side hist-write / max / scale-write chain is gone. The counter + slot stays int32-zero (float bits) between launches. ``idx`` advances + host-side each step; ``margin`` is fixed by the recipe. """ - __slots__ = ("recipe", "hist", "idx", "scale", "initialized") + __slots__ = ("recipe", "state", "hist", "scale", "idx", "initialized") def __init__(self, device: torch.device, recipe: FP8Recipe): self.recipe = recipe n = recipe.history_len - self.hist = torch.ones(n, device=device, dtype=torch.float32) + # [hist | scale | counter]; the counter slot must start at int 0. + self.state = torch.zeros(n + 2, device=device, dtype=torch.float32) + self.hist = self.state[:n] + self.scale = self.state[n : n + 1] self.idx = 0 - self.scale = torch.ones(1, device=device, dtype=torch.float32) self.initialized = False - def update(self, amax: torch.Tensor, fmt: str) -> None: - self.hist[self.idx] = amax.reshape(()) + def advance(self) -> None: + """Rotate to the next history slot after an in-kernel finalize.""" self.idx = (self.idx + 1) % self.hist.numel() - self.scale.copy_(self.recipe.scale_from_history(self.hist, fmt)) def seed(self, t: torch.Tensor, fmt: str) -> None: amax = t.abs().amax().to(torch.float32).clamp_min(1e-12) @@ -235,8 +241,9 @@ def fp8_linear_forward(x: torch.Tensor, w: torch.Tensor, bias=None): """Scaled fp8 linear forward (called from the aten::linear impl). Pure FP8 path for both recipes: quantize x/w with the active scales, run - the pre-quantized GEMM, and feed the freshly measured amax back into the - delayed-scaling ring (dynamic scaling measures the current amax itself). + the pre-quantized GEMM. With delayed scaling the rings finalize inside + the quantize kernels (amax folded into the window, next step's scale + published on device); dynamic scaling measures the current amax itself. """ if bias is None: bias = torch.empty(0, device=x.device, dtype=x.dtype) @@ -246,17 +253,34 @@ def fp8_linear_forward(x: torch.Tensor, w: torch.Tensor, bias=None): meta = None sx = _dynamic_scale(x.reshape(-1, w.size(1)), state.recipe, fmt) sw = _dynamic_scale(w, state.recipe, fmt) + out, amax_x, amax_w = linear_forward_fp8(x, w, bias, sx, sw, fmt) else: meta = state.get_weight_meta(w) if not meta.w.initialized: meta.w.seed(w, fmt) if not meta.x.initialized: meta.x.seed(x, fmt) - sx, sw = meta.x.scale, meta.w.scale - out, amax_x, amax_w = linear_forward_fp8(x, w, bias, sx, sw, fmt) - if meta is not None: - meta.x.update(amax_x, fmt) - meta.w.update(amax_w, fmt) + # In-kernel ring finalization: the kernels write hist[idx] and the + # next scale; idx rotates host-side (the device counter self-rearms). + w_is_fp8 = w.dtype != torch.bfloat16 + out, amax_x, amax_w = linear_forward_fp8( + x, + w, + bias, + meta.x.scale, + meta.w.scale, + fmt, + None, + meta.x.state, + meta.x.idx, + state.recipe.margin, + None if w_is_fp8 else meta.w.state, + meta.w.idx, + state.recipe.margin, + ) + meta.x.advance() + if not w_is_fp8: + meta.w.advance() return out @@ -290,18 +314,29 @@ class _LinearFp8(torch.autograd.Function): sg = _dynamic_scale(g, ctx.recipe, fmt) sw = _dynamic_scale(w, ctx.recipe, fmt) sx = _dynamic_scale(x, ctx.recipe, fmt) + grad_x, grad_w, grad_b, amax_g = linear_backward_fp8( + g, x, w, list(ctx.needs_input_grad), sg, sw, sx, fmt + ) else: meta = ctx.meta if not meta.g.initialized: meta.g.seed(g, fmt) - sg, sw, sx = meta.g.scale, meta.w.scale, meta.x.scale - masks = list(ctx.needs_input_grad) - grad_x, grad_w, grad_b, amax_g = linear_backward_fp8( - g, x, w, masks, sg, sw, sx, fmt - ) - if not ctx.is_dynamic: - ctx.meta.g.update(amax_g, fmt) - return grad_x, grad_w, grad_b if masks[2] else None + # The g quantize kernel finalizes the gradient's ring in-kernel. + grad_x, grad_w, grad_b, amax_g = linear_backward_fp8( + g, + x, + w, + list(ctx.needs_input_grad), + meta.g.scale, + meta.w.scale, + meta.x.scale, + fmt, + meta.g.state, + meta.g.idx, + ctx.recipe.margin, + ) + meta.g.advance() + return grad_x, grad_w, grad_b if ctx.needs_input_grad[2] else None # --------------------------------------------------------------------------- diff --git a/astrai/extension/ops/fp8.py b/astrai/extension/ops/fp8.py index 8024c44..9966525 100644 --- a/astrai/extension/ops/fp8.py +++ b/astrai/extension/ops/fp8.py @@ -138,7 +138,21 @@ def mm_fp8( return fp8_gemm(a, b, sa, sb, int(out_dtype == "e4m3"), out_scale) -def linear_forward_fp8(x, w, bias, sx, sw, fmt: str = "e4m3", bias_scale=None): +def linear_forward_fp8( + x, + w, + bias, + sx, + sw, + fmt: str = "e4m3", + bias_scale=None, + x_ring=None, + x_ring_idx: int = 0, + x_ring_margin: int = 0, + w_ring=None, + w_ring_idx: int = 0, + w_ring_margin: int = 0, +): """Pure FP8 linear forward: quantize x/w to ``fmt``, pre-quantized GEMM. Returns ``(out, amax_x, amax_w)``. ``bias`` may be ``None``. For static @@ -146,6 +160,10 @@ def linear_forward_fp8(x, w, bias, sx, sw, fmt: str = "e4m3", bias_scale=None): (produced by :func:`quantize_bf16` with their scales as ``sw`` / ``bias_scale``); a pre-quantized ``bias`` requires ``bias_scale``, and its ``amax_w`` comes back 0. The bias is fused into the GEMM epilogue. + ``x_ring`` / ``w_ring`` (delayed scaling) are ``[hist | scale | counter]`` + float32 buffers the quantize kernels finalize in-kernel: the measured + amax lands in ``hist[idx]`` and the next step's scale is published on + device, replacing the eager hist/max/scale update chain. """ fmt8 = _fmt_dtype(fmt) if x.dtype != torch.bfloat16 or w.dtype not in (torch.bfloat16, fmt8): @@ -155,16 +173,42 @@ def linear_forward_fp8(x, w, bias, sx, sw, fmt: str = "e4m3", bias_scale=None): if bias is None: bias = torch.empty(0, device=x.device, dtype=x.dtype) return get_module("fp8_ops").linear_forward_fp8( - x, w, bias, sx, sw, _fmt_int(fmt), bias_scale + x, + w, + bias, + sx, + sw, + _fmt_int(fmt), + bias_scale, + x_ring, + x_ring_idx, + x_ring_margin, + w_ring, + w_ring_idx, + w_ring_margin, ) -def linear_backward_fp8(g, x, w, masks, sg, sw, sx, fmt: str = "e5m2"): +def linear_backward_fp8( + g, + x, + w, + masks, + sg, + sw, + sx, + fmt: str = "e5m2", + g_ring=None, + g_ring_idx: int = 0, + g_ring_margin: int = 0, +): """FP8 linear backward; returns ``(grad_input, grad_weight, grad_bias, amax_g)``. The gradient (and the transposed w/x operands) are quantized to ``fmt`` (default E5M2 — larger dynamic range for gradients) and the two GEMMs run as FP8 tensor-core products sharing a single gradient quantization. + ``g_ring`` (delayed scaling) is a ``[hist | scale | counter]`` buffer the + g quantize kernel finalizes in-kernel (see :func:`linear_forward_fp8`). """ if not ( g.dtype == torch.bfloat16 @@ -175,5 +219,15 @@ def linear_backward_fp8(g, x, w, masks, sg, sw, sx, fmt: str = "e5m2"): f"fp8 backward requires bf16 inputs, got {g.dtype}/{x.dtype}/{w.dtype}" ) return get_module("fp8_ops").linear_backward_fp8( - g, x, w, list(masks), sg, sw, sx, _fmt_int(fmt) + g, + x, + w, + list(masks), + sg, + sw, + sx, + _fmt_int(fmt), + g_ring, + g_ring_idx, + g_ring_margin, ) diff --git a/csrc/kernels/fp8/common.h b/csrc/kernels/fp8/common.h index 15e9d4a..e2b02a1 100644 --- a/csrc/kernels/fp8/common.h +++ b/csrc/kernels/fp8/common.h @@ -63,16 +63,47 @@ struct Fp8GemmTraits { static constexpr float kFp8Max = kIsE5M2 ? 57344.0f : 448.0f; }; +// Quantize-kernel parameter POD: BF16 -> FP8 with fused amax and optional +// delayed-scaling ring finalization. Separate from FP8Params so each +// operator owns exactly the fields it touches (the GEMM never reads amax / +// ring state). Same NSDMI rationale: amax / ring_state gate optional paths +// via null checks. Still an aggregate, still trivially copyable. +struct FP8QuantizeParams { + // BF16 input and FP8 output buffers; scale_a is the quantization step + // (device scalar). amax_a (may be null) is zero-initialized by the + // binding and receives the raw-domain absolute maximum. + const void* __restrict__ a_ptr = nullptr; + void* __restrict__ out_ptr = nullptr; + const float* __restrict__ scale_a = nullptr; + float* __restrict__ amax_a = nullptr; + + // Optional delayed-scaling ring finalization. ring_state packs + // [hist[ring_len] | scale | counter] with ring_len = numel - 2. When + // non-null and amax_a is set, the last-finishing block records the + // measured amax into hist[ring_idx], reduces the window and publishes + // the next step's scale (max(hist) / fp8_max / 2^ring_margin) — the + // fused replacement for the eager hist-write / max / scale-write chain, + // at zero extra launches. The counter slot is a persistent zero-armed + // int32 (float bits) electing the last block each launch. + float* ring_state = nullptr; + int ring_len = 0; + int ring_idx = 0; + int ring_margin = 0; + + // Element count (only the elementwise quantize kernel uses it). + int total = 0; +}; + // Unified GEMM parameter POD, mirroring AttentionParams: one struct flows -// through quantize / fused / pre-quantized kernels. Each kernel touches only -// the fields it needs; buffers are raw pointers packed by the torch binding. +// through the pre-quantized GEMM kernels. Each kernel touches only the +// fields it needs; buffers are raw pointers packed by the torch binding. // Pointer members default to null (same NSDMI rationale as AttentionParams: -// bias / amax / out_scale gate optional paths via null checks, so a partially +// bias / out_scale gate optional paths via null checks, so a partially // packed struct must never hold garbage non-null pointers). Still an // aggregate, still trivially copyable. struct FP8Params { - // Inputs: a/b are BF16 for the fused (quantize-in-GEMM) path, FP8 for - // the pre-quantized path. Scales are quantization steps (device scalars). + // Inputs: a/b are FP8 for the pre-quantized path. Scales are + // quantization steps (device scalars). const void* __restrict__ a_ptr = nullptr; const void* __restrict__ b_ptr = nullptr; const void* __restrict__ bias = nullptr; @@ -84,23 +115,16 @@ struct FP8Params { void* __restrict__ out_ptr = nullptr; const float* __restrict__ out_scale = nullptr; - // Fused forward extras: bias (may be null) and amax slots (may be null). - float* __restrict__ amax_a = nullptr; - float* __restrict__ amax_b = nullptr; - - // Shapes. total is only used by the elementwise quantize kernel. `int` - // covers every realistic LLM shape; the kernels promote to int64 for all - // pointer arithmetic. + // Shapes. `int` covers every realistic LLM shape; the kernels promote + // to int64 for all pointer arithmetic. int m, n, k; - // Physical leading dimensions (column count, i.e. row stride) of A and B. - // For a non-transposed operand the stride equals the contract dim; for a - // transposed operand it is the operand's own column count. The binding - // packs these so the kernel reads both buffers either naturally or - // transposed depending on the LayoutA/LayoutB tags (see gemm.cuh). + // Physical leading dimensions (column count, i.e. row stride) of A and + // B. For a non-transposed operand the stride equals the contract dim; + // for a transposed operand it is the operand's own column count. The + // binding packs these so the kernel reads both buffers either naturally + // or transposed depending on the LayoutA/LayoutB tags (see gemm.cuh). int a_ld, b_ld; - - int total; }; } // namespace fp8 diff --git a/csrc/kernels/fp8/gemm.cuh b/csrc/kernels/fp8/gemm.cuh index d9b0c54..657d243 100644 --- a/csrc/kernels/fp8/gemm.cuh +++ b/csrc/kernels/fp8/gemm.cuh @@ -70,7 +70,7 @@ __device__ __forceinline__ unsigned quantize2(unsigned pair, float inv, } template -__global__ void fp8_quantize_kernel(FP8Params p) { +__global__ void fp8_quantize_kernel(FP8QuantizeParams p) { const float inv = 1.0f / *p.scale_a; const auto* x = reinterpret_cast(p.a_ptr); void* x8 = p.out_ptr; @@ -122,6 +122,51 @@ __global__ void fp8_quantize_kernel(FP8Params p) { atomic_max_float(amax, v); } } + if (p.ring_state && amax) { + // Delayed-scaling ring finalization as a last-block epilogue (the + // CUDA threadFenceReduction pattern): the fence + counter elect the + // final block once every block's atomic_max above is visible; warp 0 + // folds the fresh amax into the window, reduces it and publishes the + // next step's scale, then re-arms the counter for the next launch. + // __fdiv_rn / ldexpf keep the scale bit-identical to the eager + // (peak / fp8_max) / 2^margin fp32 chain despite --use_fast_math. + __threadfence(); + __shared__ bool ring_last; + if (threadIdx.x == 0) + ring_last = atomicAdd(reinterpret_cast(p.ring_state + + p.ring_len + 1), + 1) == gridDim.x - 1; + __syncthreads(); + if (ring_last && threadIdx.x < 32) { + float* hist = p.ring_state; + const int lane = threadIdx.x; + float v = 0.0f; + if (lane < p.ring_len) v = hist[lane]; + if (lane == p.ring_idx) { + v = *amax; // the global amax is final now + hist[lane] = v; + } + // Windows longer than one warp (atypical) fold the tail. + for (int i = lane + 32; i < p.ring_len; i += 32) { + float h = hist[i]; + if (i == p.ring_idx) { + h = *amax; + hist[i] = h; + } + v = fmaxf(v, h); + } + const float peak = warp_reduce_max(v); + if (lane == 0) { + constexpr float kFmtMax = + Fmt == FP8Format::E5M2 ? 57344.0f : 448.0f; + p.ring_state[p.ring_len] = fmaxf( + ldexpf(__fdiv_rn(peak, kFmtMax), -p.ring_margin), 1e-12f); + __threadfence(); + // Re-arm the counter (0.0f bits == int32 0). + p.ring_state[p.ring_len + 1] = 0.0f; + } + } + } } // Swizzled address inside a flat [rows * K] staging tile: the 16-byte chunk @@ -539,7 +584,7 @@ __global__ void // --------------------------------------------------------------------------- template -void launch_fp8_quantize(const FP8Params& p, cudaStream_t stream) { +void launch_fp8_quantize(const FP8QuantizeParams& p, cudaStream_t stream) { constexpr int kThreads = 256; // One block per 256 vectors (8 elements each); at least one block so the // scalar tail of a tiny / misaligned tensor is still covered. diff --git a/csrc/kernels/fp8/ops.cu b/csrc/kernels/fp8/ops.cu index 5a5ebbb..cdb8b2c 100644 --- a/csrc/kernels/fp8/ops.cu +++ b/csrc/kernels/fp8/ops.cu @@ -71,30 +71,34 @@ void pack_gemm_params(FP8Params& p, const void* a, const void* b, void* out, p.out_scale = out_scale ? out_scale->data_ptr() : nullptr; p.bias = bias; p.bias_scale = bias_scale ? bias_scale->data_ptr() : nullptr; - p.amax_a = nullptr; - p.amax_b = nullptr; p.m = static_cast(m); p.n = static_cast(n); p.k = static_cast(k); p.a_ld = static_cast(a_ld); p.b_ld = static_cast(b_ld); - p.total = 0; } -void pack_quantize_params(FP8Params& p, const void* x, void* x8, +// Pack the quantize params, optionally wiring the delayed-scaling ring. +// ring (may be null) packs [hist[len] | scale | counter]; len/margin come +// from the active recipe and idx is the caller's slot for this step. +void pack_quantize_params(FP8QuantizeParams& p, const void* x, void* x8, const torch::Tensor& scale, torch::Tensor* amax, - int64_t total) { + const torch::Tensor* ring, int64_t ring_idx, + int64_t ring_margin, int64_t total) { p.a_ptr = x; - p.b_ptr = nullptr; p.out_ptr = x8; p.scale_a = scale.data_ptr(); - p.scale_b = nullptr; - p.out_scale = nullptr; - p.bias = nullptr; p.amax_a = amax ? amax->data_ptr() : nullptr; - p.amax_b = nullptr; - p.m = p.n = p.k = 0; - p.a_ld = p.b_ld = 0; + if (ring && ring->defined()) { + TORCH_CHECK(ring->is_cuda() && ring->scalar_type() == torch::kFloat32 && + ring->numel() >= 3 && ring->is_contiguous(), + "ring must be a contiguous CUDA float32 tensor packing " + "[hist | scale | counter]"); + p.ring_state = ring->data_ptr(); + p.ring_len = static_cast(ring->numel() - 2); + p.ring_idx = static_cast(ring_idx); + p.ring_margin = static_cast(ring_margin); + } p.total = static_cast(total); } @@ -152,11 +156,11 @@ std::tuple quantize_bf16(torch::Tensor x, auto x_c = x.contiguous(); auto x8 = torch::empty_like( x_c, x_c.options().dtype(fmt ? torch::kFloat8_e5m2 - : torch::kFloat8_e4m3fn)); + : torch::kFloat8_e4m3fn)); auto amax = torch::zeros({1}, x_c.options().dtype(torch::kFloat32)); - FP8Params p; + FP8QuantizeParams p; pack_quantize_params(p, x_c.data_ptr(), x8.data_ptr(), scale, &amax, - x_c.numel()); + nullptr, 0, 0, x_c.numel()); if (fmt) { launch_fp8_quantize(p, stream.stream()); } else { @@ -228,14 +232,19 @@ torch::Tensor mm_fp8(torch::Tensor a, torch::Tensor b, torch::Tensor sa, std::tuple linear_forward_fp8( torch::Tensor x, torch::Tensor w, torch::Tensor bias, torch::Tensor sx, - torch::Tensor sw, int64_t fmt, - c10::optional bias_scale) { + torch::Tensor sw, int64_t fmt, c10::optional bias_scale, + c10::optional x_ring, int64_t x_ring_idx, + int64_t x_ring_margin, c10::optional w_ring, + int64_t w_ring_idx, int64_t w_ring_margin) { // Pure FP8 forward: quantize x/w (fmt: 0 = E4M3, 1 = E5M2), then the // pre-quantized GEMM; the dequantized BF16 output gets the bias added. // amax_x / amax_w come from the quantize kernels (zero-initialized here; // a pre-quantized w reports amax_w = 0 — nothing to feed a delayed ring). // w may itself be pre-quantized fp8 storage matching fmt (static // inference weights): the weight quantize is skipped, amax_w stays 0. + // When x_ring / w_ring are given (delayed scaling), the quantize kernels + // finalize them in-kernel: the returned amax is already folded into the + // ring window and the next step's scale is published on device. TORCH_CHECK(x.is_cuda() && w.is_cuda(), "CUDA tensors required"); const auto f8opt = fmt ? torch::kFloat8_e5m2 : torch::kFloat8_e4m3fn; const bool w_prequant = w.scalar_type() == f8opt; @@ -272,9 +281,12 @@ std::tuple linear_forward_fp8( auto out = torch::empty({m, n}, x_c.options()); auto quantize = [&](const torch::Tensor& src, torch::Tensor& dst, - const torch::Tensor& scale, torch::Tensor* amax) { - FP8Params qp; + const torch::Tensor& scale, torch::Tensor* amax, + const c10::optional& ring, + int64_t ring_idx, int64_t ring_margin) { + FP8QuantizeParams qp; pack_quantize_params(qp, src.data_ptr(), dst.data_ptr(), scale, amax, + ring ? &*ring : nullptr, ring_idx, ring_margin, src.numel()); if (fmt) { launch_fp8_quantize(qp, stream.stream()); @@ -282,13 +294,14 @@ std::tuple linear_forward_fp8( launch_fp8_quantize(qp, stream.stream()); } }; - quantize(x_c, x8, sx, &amax_x); + quantize(x_c, x8, sx, &amax_x, x_ring, x_ring_idx, x_ring_margin); // Static inference weights arrive pre-quantized (w8 storage + its scale); // only freshly-loaded bf16 weights quantize here. torch::Tensor w8 = w_prequant ? w_c : torch::empty({n, k}, x_c.options().dtype(f8opt)); - if (!w_prequant) quantize(w_c, w8, sw, &amax_w); + if (!w_prequant) + quantize(w_c, w8, sw, &amax_w, w_ring, w_ring_idx, w_ring_margin); FP8Params p; // Forward is the NT layout: A = x8 [M,K] (a_ld = k), B = w8 [N,K] @@ -315,10 +328,16 @@ std::tuple linear_forward_fp8( std::tuple linear_backward_fp8(torch::Tensor g, torch::Tensor x, torch::Tensor w, std::vector masks, torch::Tensor sg, - torch::Tensor sw, torch::Tensor sx, int64_t fmt) { + torch::Tensor sw, torch::Tensor sx, int64_t fmt, + c10::optional g_ring, int64_t g_ring_idx, + int64_t g_ring_margin) { // Pre-quantized FP8 backward: grad is quantized once (E4M3 or E5M2 per // `fmt`), then dX / dW run as FP8 tensor-core GEMMs sharing g8. - // Returns (grad_input, grad_weight, grad_bias, amax_g). + // Returns (grad_input, grad_weight, grad_bias, amax_g). With g_ring + // (delayed scaling), the g quantize kernel finalizes the ring in-kernel + // (amax folded into the window, next step's scale published on device); + // the w/x quantizes for dX / dW never touch rings — each operand's ring + // is finalized exactly once per step (by the forward or this kernel). TORCH_CHECK(g.is_cuda() && x.is_cuda() && w.is_cuda(), "CUDA tensors required"); TORCH_CHECK(g.scalar_type() == torch::kBFloat16 && x.scalar_type() == torch::kBFloat16 && @@ -346,9 +365,12 @@ linear_backward_fp8(torch::Tensor g, torch::Tensor x, torch::Tensor w, : g.options().dtype(torch::kFloat8_e4m3fn); auto quantize = [&](const torch::Tensor& src, torch::Tensor& dst, - const torch::Tensor& scale, torch::Tensor* amax) { - FP8Params qp; + const torch::Tensor& scale, torch::Tensor* amax, + const c10::optional& ring, + int64_t ring_idx, int64_t ring_margin) { + FP8QuantizeParams qp; pack_quantize_params(qp, src.data_ptr(), dst.data_ptr(), scale, amax, + ring ? &*ring : nullptr, ring_idx, ring_margin, src.numel()); if (fmt) { launch_fp8_quantize(qp, stream.stream()); @@ -375,13 +397,13 @@ linear_backward_fp8(torch::Tensor g, torch::Tensor x, torch::Tensor w, torch::Tensor g8; if (masks[0] || masks[1]) { g8 = torch::empty({m, n}, f8opt); - quantize(g_c, g8, sg, &amax_g); + quantize(g_c, g8, sg, &amax_g, g_ring, g_ring_idx, g_ring_margin); } // dX = g @ w: A = g8 [M,N] (contract over N), B = w8 [N,K] read transposed // (b[p*b_ld + n] = w[p,n]); out = [M,K], a_ld = N, b_ld = K, contract = N. if (masks[0]) { auto w8 = torch::empty({n, k}, f8opt); - quantize(w_c, w8, sw, nullptr); + quantize(w_c, w8, sw, nullptr, c10::nullopt, 0, 0); auto grad_input_2d = grad_input.reshape({m, k}); FP8Params gp; pack_gemm_params(gp, g8.data_ptr(), w8.data_ptr(), @@ -394,7 +416,7 @@ linear_backward_fp8(torch::Tensor g, torch::Tensor x, torch::Tensor w, // b_ld = K, contract = M. if (masks[1]) { auto x8 = torch::empty({m, k}, f8opt); - quantize(x_c, x8, sx, nullptr); + quantize(x_c, x8, sx, nullptr, c10::nullopt, 0, 0); FP8Params gp; pack_gemm_params(gp, g8.data_ptr(), x8.data_ptr(), grad_weight.data_ptr(), sg, sx, nullptr, nullptr, @@ -423,12 +445,21 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { m.def("linear_forward_fp8", &linear_forward_fp8, py::arg("x"), py::arg("w"), py::arg("bias"), py::arg("sx"), py::arg("sw"), py::arg("fmt") = 0, py::arg("bias_scale") = py::none(), + py::arg("x_ring") = py::none(), py::arg("x_ring_idx") = 0, + py::arg("x_ring_margin") = 0, py::arg("w_ring") = py::none(), + py::arg("w_ring_idx") = 0, py::arg("w_ring_margin") = 0, "Pure FP8 linear forward: quantize x/w, pre-quantized GEMM with the " "bias fused into the epilogue; w and bias may be pre-quantized fp8 " "matching fmt (static inference path; fp8 bias requires bias_scale);" - " returns (out, amax_x, amax_w)"); + " x_ring/w_ring optionally finalize a delayed-scaling ring " + "([hist | scale | counter] float32 buffer) in-kernel; returns " + "(out, amax_x, amax_w)"); m.def("linear_backward_fp8", &linear_backward_fp8, py::arg("g"), py::arg("x"), py::arg("w"), py::arg("masks"), py::arg("sg"), py::arg("sw"), py::arg("sx"), py::arg("fmt"), - "FP8 linear backward; returns (grad_input, grad_weight, grad_bias, amax_g)"); + py::arg("g_ring") = py::none(), py::arg("g_ring_idx") = 0, + py::arg("g_ring_margin") = 0, + "FP8 linear backward; g_ring optionally finalizes the gradient's " + "delayed-scaling ring in-kernel; returns (grad_input, grad_weight, " + "grad_bias, amax_g)"); } diff --git a/tests/extension/test_fp8_mma.py b/tests/extension/test_fp8_mma.py index b4a9ca8..1437bff 100644 --- a/tests/extension/test_fp8_mma.py +++ b/tests/extension/test_fp8_mma.py @@ -86,6 +86,51 @@ def test_quantize_bf16_e5m2_format(): torch.testing.assert_close(amax, x.abs().amax().float().reshape(1)) +@skip_no_fp8 +def test_quantize_ring_in_kernel_finalize(): + """The quantize kernel finalizes the delayed-scaling ring in-kernel: the + measured amax lands in hist[idx], the window reduces to the next step's + scale on device, and the counter re-arms for the next launch.""" + from astrai.extension.fp8 import _ScaleRing + + torch.manual_seed(21) + dev = torch.device("cuda") + ring = _ScaleRing(dev, DelayedScaling(history_len=4, margin=0)) + x0 = torch.randn(256, 256, device=dev, dtype=torch.bfloat16) + w = torch.randn(256, 256, device=dev, dtype=torch.bfloat16) + sw = torch.tensor([1.0], device=dev) + ring.seed(x0, "e4m3") + hist0 = ring.hist.clone() + + # Step over three fresh tensors: each launch folds its amax into + # hist[idx] and publishes max(hist)/448 as the next scale. + idx = 0 + for _ in range(3): + x = torch.randn(256, 256, device=dev, dtype=torch.bfloat16) * (2.0 + 4.0 * _) + _ = linear_forward_fp8( + x, + w, + None, + ring.scale, + sw, + "e4m3", + None, + ring.state, + idx, + 0, + ) + torch.cuda.synchronize() + expected_hist = hist0.clone() + expected_hist[idx] = x.abs().amax().float() + torch.testing.assert_close(ring.hist, expected_hist) + expected_scale = (expected_hist.max() / 448.0).reshape(1) + torch.testing.assert_close(ring.scale, expected_scale, rtol=1e-6, atol=1e-12) + # counter re-armed to int32 zero + assert ring.state[-1].view(torch.int32).item() == 0 + hist0 = expected_hist.clone() + idx = (idx + 1) % 4 + + @skip_no_fp8 def test_fp8_linear_forward_and_backward(): torch.manual_seed(7) @@ -196,9 +241,9 @@ def test_fp8_linear_backward_outside_autocast(): calls = {"bwd": 0} orig = f8mod.linear_backward_fp8 - def spy(g, xx, ww, masks, sg, sw, sx, fmt="e5m2"): + def spy(*args, **kwargs): calls["bwd"] += 1 - return orig(g, xx, ww, masks, sg, sw, sx, fmt) + return orig(*args, **kwargs) f8mod.linear_backward_fp8 = spy try: @@ -331,14 +376,20 @@ def test_fp8_autocast_context(): def test_fp8_tensor_meta_delayed_update(): - """Meta seeds from data and refreshes the scale from the amax ring.""" + """Meta seeds from data; hist/scale are packed views of one state buffer.""" meta = FP8TensorMeta(torch.device("cpu"), DelayedScaling(history_len=4, margin=0)) w = torch.randn(8, 8) meta.w.seed(w, "e4m3") assert meta.w.initialized torch.testing.assert_close(meta.w.scale, (w.abs().amax() / 448.0).reshape(1)) - meta.w.update(torch.tensor([4.0]), "e4m3") - torch.testing.assert_close(meta.w.scale, torch.tensor(4.0 / 448.0).reshape(1)) + # [hist | scale | counter] packing: views alias the single state buffer. + assert meta.w.state.numel() == 4 + 2 + assert meta.w.hist.data_ptr() == meta.w.state.data_ptr() + assert meta.w.scale.data_ptr() == meta.w.state[4:].data_ptr() + # counter slot stays int32-zero (float bits) between launches + assert meta.w.state[-1].view(torch.int32).item() == 0 + meta.w.advance() + assert meta.w.idx == 1 def test_quantize_bf16_cpu_fallback():