diff --git a/astrai/extension/ops/fp8.py b/astrai/extension/ops/fp8.py index 739051c..8024c44 100644 --- a/astrai/extension/ops/fp8.py +++ b/astrai/extension/ops/fp8.py @@ -138,18 +138,25 @@ 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"): +def linear_forward_fp8(x, w, bias, sx, sw, fmt: str = "e4m3", bias_scale=None): """Pure FP8 linear forward: quantize x/w to ``fmt``, pre-quantized GEMM. - Returns ``(out, amax_x, amax_w)``. ``bias`` may be ``None``. Both - operands share the same FP8 format (E4M3 by default; E5M2 for a - range-first configuration). + Returns ``(out, amax_x, amax_w)``. ``bias`` may be ``None``. For static + fp8 inference, ``w`` and ``bias`` may arrive pre-quantized to ``fmt`` + (produced by :func:`quantize_bf16` with their scales as ``sw`` / + ``bias_scale``); a pre-quantized ``bias`` requires ``bias_scale``, and + its ``amax_w`` comes back 0. The bias is fused into the GEMM epilogue. """ - if not (x.dtype == torch.bfloat16 and w.dtype == torch.bfloat16): - raise TypeError(f"fp8 forward requires bf16 inputs, got {x.dtype}/{w.dtype}") + fmt8 = _fmt_dtype(fmt) + if x.dtype != torch.bfloat16 or w.dtype not in (torch.bfloat16, fmt8): + raise TypeError( + f"fp8 forward requires bf16 x and bf16-or-{fmt} w, got {x.dtype}/{w.dtype}" + ) if bias is None: bias = torch.empty(0, device=x.device, dtype=x.dtype) - return get_module("fp8_ops").linear_forward_fp8(x, w, bias, sx, sw, _fmt_int(fmt)) + return get_module("fp8_ops").linear_forward_fp8( + x, w, bias, sx, sw, _fmt_int(fmt), bias_scale + ) def linear_backward_fp8(g, x, w, masks, sg, sw, sx, fmt: str = "e5m2"): diff --git a/csrc/kernels/fp8/common.h b/csrc/kernels/fp8/common.h index bc8ac04..15e9d4a 100644 --- a/csrc/kernels/fp8/common.h +++ b/csrc/kernels/fp8/common.h @@ -75,16 +75,16 @@ struct FP8Params { // 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; const float* __restrict__ scale_a = nullptr; const float* __restrict__ scale_b = nullptr; - + const float* __restrict__ bias_scale = nullptr; // Output: BF16 or FP8 (E4M3). out_scale is the output quantization step // (FP8 output only). void* __restrict__ out_ptr = nullptr; const float* __restrict__ out_scale = nullptr; // Fused forward extras: bias (may be null) and amax slots (may be null). - const __nv_bfloat16* __restrict__ bias = nullptr; float* __restrict__ amax_a = nullptr; float* __restrict__ amax_b = nullptr; diff --git a/csrc/kernels/fp8/gemm.cuh b/csrc/kernels/fp8/gemm.cuh index e905a32..f24b155 100644 --- a/csrc/kernels/fp8/gemm.cuh +++ b/csrc/kernels/fp8/gemm.cuh @@ -272,8 +272,7 @@ __device__ __forceinline__ unsigned frag_addr(const T8* tile, int row, // exists for small-M calls: m <= 64 wastes half of every 128-row CTA, so the // launcher dispatches to it there (see launch_fp8_gemm). template -__global__ void __launch_bounds__( - (Traits::kBlockM / 64) * (Traits::kBlockN / 32) * 32, 2) +__global__ void __launch_bounds__((Traits::kBlockM / 64) * (Traits::kBlockN / 32) * 32, 2) fp8_gemm_kernel(FP8Params p) { using T8 = std::conditional_t; constexpr int kBlockM = Traits::kBlockM; @@ -453,36 +452,48 @@ fp8_gemm_kernel(FP8Params p) { } const float output_scale = sa * sb; - const float o8_scale = OutFp8 ? output_scale * *p.out_scale : 0.0f; + // Fused bias: BF16 raw values, or FP8 storage dequantized by its own + // scale (bias_scale != null selects the FP8 path; the format follows the + // kernel's Traits). Added in real units after the operand dequantization + // and before any output quantization. + const auto* bias16 = static_cast(p.bias); + const auto* bias8 = static_cast(p.bias); + auto bias_val = [&](int64_t col) -> float { + if (p.bias == nullptr || col >= n) return 0.0f; + if (p.bias_scale == nullptr) return __bfloat162float(bias16[col]); + return __half2float(__half(bias8[col])) * *p.bias_scale; + }; #pragma unroll for (int nt = 0; nt < 4; ++nt) { const int64_t col = output_col + nt * 8; + const float b0 = bias_val(col); + const float b1 = bias_val(col + 1); // Per-row store: FP8 packs two adjacent columns into one 16-bit // write, BF16 into one 32-bit __nv_bfloat162 (single cvt+pack // instruction); boundary or unaligned columns fall back to scalar // converts so a pack never crosses the row edge or misaligns. auto store_out = [&](int64_t row, float v0, float v1) { if (row >= m) return; + const float r0 = v0 * output_scale + b0; + const float r1 = v1 * output_scale + b1; if constexpr (OutFp8) { if (col + 1 < n) { *reinterpret_cast(out_fp8 + row * n + col) = static_cast(__nv_cvt_float2_to_fp8x2( - make_float2(v0 * o8_scale, v1 * o8_scale), + make_float2(r0 * *p.out_scale, r1 * *p.out_scale), __NV_SATFINITE, __NV_E4M3)); } else { - out_fp8[row * n + col] = __nv_fp8_e4m3(v0 * o8_scale); + out_fp8[row * n + col] = __nv_fp8_e4m3(r0 * *p.out_scale); } } else { auto* dst = out_bf16 + row * n + col; if (col + 1 < n && (reinterpret_cast(dst) & 3) == 0) { *reinterpret_cast<__nv_bfloat162*>(dst) = - __floats2bfloat162_rn(v0 * output_scale, - v1 * output_scale); + __floats2bfloat162_rn(r0, r1); } else { - dst[0] = __float2bfloat16(v0 * output_scale); - if (col + 1 < n) - dst[1] = __float2bfloat16(v1 * output_scale); + dst[0] = __float2bfloat16(r0); + if (col + 1 < n) dst[1] = __float2bfloat16(r1); } } }; diff --git a/csrc/kernels/fp8/ops.cu b/csrc/kernels/fp8/ops.cu index 2f462f1..5a5ebbb 100644 --- a/csrc/kernels/fp8/ops.cu +++ b/csrc/kernels/fp8/ops.cu @@ -60,7 +60,8 @@ void check_scale(const torch::Tensor& scale, const torch::Tensor& input, void pack_gemm_params(FP8Params& p, const void* a, const void* b, void* out, const torch::Tensor& sa, const torch::Tensor& sb, - const torch::Tensor* out_scale, int64_t m, int64_t n, + const torch::Tensor* out_scale, const void* bias, + const torch::Tensor* bias_scale, int64_t m, int64_t n, int64_t k, int64_t a_ld, int64_t b_ld) { p.a_ptr = a; p.b_ptr = b; @@ -68,7 +69,8 @@ void pack_gemm_params(FP8Params& p, const void* a, const void* b, void* out, p.scale_a = sa.data_ptr(); p.scale_b = sb.data_ptr(); p.out_scale = out_scale ? out_scale->data_ptr() : nullptr; - p.bias = 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); @@ -214,7 +216,8 @@ torch::Tensor mm_fp8(torch::Tensor a, torch::Tensor b, torch::Tensor sa, : a_c.options().dtype(torch::kBFloat16)); FP8Params p; pack_gemm_params(p, a_c.data_ptr(), b_c.data_ptr(), out.data_ptr(), sa, sb, - out_fp8 ? &os : nullptr, m, n, k, a_ld, b_ld); + out_fp8 ? &os : nullptr, nullptr, nullptr, m, n, k, a_ld, + b_ld); if (a.scalar_type() == torch::kFloat8_e4m3fn) dispatch_gemm(p, stream.stream(), out_fp8, ta, tb); else @@ -225,14 +228,21 @@ 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) { + torch::Tensor sw, int64_t fmt, + c10::optional bias_scale) { // 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). + // 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. TORCH_CHECK(x.is_cuda() && w.is_cuda(), "CUDA tensors required"); - TORCH_CHECK(x.scalar_type() == torch::kBFloat16 && - w.scalar_type() == torch::kBFloat16, - "x and w must be bf16"); + const auto f8opt = fmt ? torch::kFloat8_e5m2 : torch::kFloat8_e4m3fn; + const bool w_prequant = w.scalar_type() == f8opt; + TORCH_CHECK( + x.scalar_type() == torch::kBFloat16 && + (w.scalar_type() == torch::kBFloat16 || w_prequant), + "x must be bf16; w must be bf16 or pre-quantized fp8 matching fmt"); TORCH_CHECK(x.device() == w.device(), "x and w must be on the same device"); check_scale(sx, x, "sx"); check_scale(sw, x, "sw"); @@ -245,15 +255,18 @@ std::tuple linear_forward_fp8( int64_t m = x_c.size(0), k = x_c.size(1), n = w_c.size(0); TORCH_CHECK(w_c.dim() == 2 && w_c.size(1) == k, "inner dim mismatch"); const bool has_bias = bias.defined() && bias.numel() > 0; + const bool b_prequant = has_bias && bias.scalar_type() == f8opt; if (has_bias) { TORCH_CHECK(bias.is_cuda() && bias.device() == x.device() && - bias.scalar_type() == torch::kBFloat16 && - bias.numel() == n, - "bias must be CUDA bf16 with shape [N]"); + bias.numel() == n && + (bias.scalar_type() == torch::kBFloat16 || b_prequant), + "bias must be CUDA bf16 or pre-quantized fp8 matching fmt, " + "with shape [N]"); + TORCH_CHECK(b_prequant == bias_scale.has_value(), + "fp8 bias requires bias_scale (and bf16 bias takes none)"); + if (b_prequant) check_scale(*bias_scale, x, "bias_scale"); } - const auto f8opt = fmt ? torch::kFloat8_e5m2 : torch::kFloat8_e4m3fn; auto x8 = torch::empty({m, k}, x_c.options().dtype(f8opt)); - auto w8 = torch::empty({n, k}, x_c.options().dtype(f8opt)); auto amax_x = torch::zeros({1}, x.options().dtype(torch::kFloat32)); auto amax_w = torch::zeros({1}, x.options().dtype(torch::kFloat32)); auto out = torch::empty({m, n}, x_c.options()); @@ -270,13 +283,21 @@ std::tuple linear_forward_fp8( } }; quantize(x_c, x8, sx, &amax_x); - quantize(w_c, w8, sw, &amax_w); + // 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); FP8Params p; // Forward is the NT layout: A = x8 [M,K] (a_ld = k), B = w8 [N,K] - // (b_ld = k), out = x @ w^T. No operand transposes needed. + // (b_ld = k), out = x @ w^T. The bias is fused into the epilogue (bf16 + // raw, or fp8 + bias_scale on the static path). + auto bias_c = has_bias ? bias.contiguous() : bias; pack_gemm_params(p, x8.data_ptr(), w8.data_ptr(), out.data_ptr(), sx, sw, - nullptr, m, n, k, k, k); + nullptr, has_bias ? bias_c.data_ptr() : nullptr, + b_prequant ? &*bias_scale : nullptr, m, n, k, k, k); if (fmt) { launch_fp8_gemm( p, stream.stream()); @@ -288,9 +309,7 @@ std::tuple linear_forward_fp8( std::vector shape(x.sizes().begin(), x.sizes().end() - 1); shape.push_back(n); - auto out_r = out.reshape(shape); - if (has_bias) out_r = out_r + bias; - return {out_r, amax_x, amax_w}; + return {out.reshape(shape), amax_x, amax_w}; } std::tuple @@ -366,8 +385,8 @@ linear_backward_fp8(torch::Tensor g, torch::Tensor x, torch::Tensor w, auto grad_input_2d = grad_input.reshape({m, k}); FP8Params gp; pack_gemm_params(gp, g8.data_ptr(), w8.data_ptr(), - grad_input_2d.data_ptr(), sg, sw, nullptr, m, k, n, n, - k); + grad_input_2d.data_ptr(), sg, sw, nullptr, nullptr, + nullptr, m, k, n, n, k); run_bwd_gemm(gp, false, false); } // dW = g^T @ x: A = g8 [M,N] read transposed (a[p*a_ld + m] = g[p,m]), B = @@ -378,7 +397,8 @@ linear_backward_fp8(torch::Tensor g, torch::Tensor x, torch::Tensor w, quantize(x_c, x8, sx, nullptr); FP8Params gp; pack_gemm_params(gp, g8.data_ptr(), x8.data_ptr(), - grad_weight.data_ptr(), sg, sx, nullptr, n, k, m, n, k); + grad_weight.data_ptr(), sg, sx, nullptr, nullptr, + nullptr, n, k, m, n, k); run_bwd_gemm(gp, true, false); } if (!masks[0] && !masks[1]) { @@ -402,9 +422,11 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { "the operand layout (default 0/0 = a@b)"); 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, - "Pure FP8 linear forward: quantize x/w, pre-quantized GEMM; " - "returns (out, amax_x, amax_w)"); + py::arg("fmt") = 0, py::arg("bias_scale") = py::none(), + "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)"); 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"), diff --git a/tests/extension/test_fp8_mma.py b/tests/extension/test_fp8_mma.py index f36ba8a..9691a5b 100644 --- a/tests/extension/test_fp8_mma.py +++ b/tests/extension/test_fp8_mma.py @@ -146,6 +146,36 @@ def test_linear_backward_e5m2_gradients(): torch.testing.assert_close(amax_g, grad.abs().amax().float().reshape(1)) +@skip_no_fp8 +def test_fp8_linear_static_fp8_weight_and_bias(): + """Static fp8 inference: pre-quantized w8/b8 + their scales take the GEMM + directly (no weight quantize, amax_w = 0); the bias is fused in the + epilogue (bf16 and fp8 bias share the fused path).""" + torch.manual_seed(9) + m, n, k = 67, 45, 129 + x = torch.randn(m, k, device="cuda", dtype=torch.bfloat16) + weight = torch.randn(n, k, device="cuda", dtype=torch.bfloat16) * 0.5 + bias = torch.randn(n, device="cuda", dtype=torch.bfloat16) * 0.5 + sx, sw, sb = _scale(x), _scale(weight), _scale(bias) + + w8, _ = quantize_bf16(weight, sw, "e4m3") + b8, _ = quantize_bf16(bias, sb, "e4m3") + out, amax_x, amax_w = linear_forward_fp8(x, w8, b8, sx, sw, "e4m3", sb) + + qx = _quantize(x, sx) + qw = _quantize(weight, sw) + qb = _quantize(bias, sb) + expected = (qx @ qw.t() * sx * sw + qb * sb).to(torch.bfloat16) + torch.testing.assert_close(out, expected, atol=0.125, rtol=0.01) + torch.testing.assert_close(amax_x, x.abs().amax().float().reshape(1)) + assert amax_w.item() == 0.0 # nothing measured on the static path + + # bf16 bias stays bf16 on the same fused-epilogue path + out_bf16bias, _, _ = linear_forward_fp8(x, w8, bias, sx, sw, "e4m3") + expected_b = (qx @ qw.t() * sx * sw + bias.float()).to(torch.bfloat16) + torch.testing.assert_close(out_bf16bias, expected_b, atol=0.125, rtol=0.01) + + @skip_no_fp8 def test_fp8_linear_backward_outside_autocast(): """aten::linear records an fp8 autograd node inside fp8_autocast; the