feat: static fp8 weights and bias with fused epilogue
- linear_forward_fp8 accepts pre-quantized w8 (matching fmt) and skips the weight quantize; amax_w returns 0 on that path since no bf16 values are seen - bias is now fused into the GEMM epilogue for both dtypes, replacing the separate torch-level add (one elementwise kernel per linear removed) - FP8Params.bias becomes void* with a new bias_scale slot: null scale = raw bf16 bias, non-null = fp8 storage dequantized in the epilogue after the operand scaling and before any output quantization - ops/fp8.py relaxes the w dtype check to bf16-or-fp8 and passes bias_scale through - regression test covers w8/b8, w8/bf16-bias and the amax_w = 0 contract vs an explicit quantization reference
This commit is contained in:
@@ -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"):
|
||||
|
||||
@@ -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;
|
||||
|
||||
|
||||
+21
-10
@@ -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 <typename Traits, bool OutFp8 = false, typename LayoutA = RowMajor, typename LayoutB = RowMajor>
|
||||
__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<Traits::kIsE5M2, __nv_fp8_e5m2, __nv_fp8_e4m3>;
|
||||
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<const __nv_bfloat16*>(p.bias);
|
||||
const auto* bias8 = static_cast<const T8*>(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<unsigned short*>(out_fp8 + row * n + col) =
|
||||
static_cast<unsigned short>(__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<uintptr_t>(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);
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
+47
-25
@@ -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<float>();
|
||||
p.scale_b = sb.data_ptr<float>();
|
||||
p.out_scale = out_scale ? out_scale->data_ptr<float>() : nullptr;
|
||||
p.bias = nullptr;
|
||||
p.bias = bias;
|
||||
p.bias_scale = bias_scale ? bias_scale->data_ptr<float>() : nullptr;
|
||||
p.amax_a = nullptr;
|
||||
p.amax_b = nullptr;
|
||||
p.m = static_cast<int>(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<FP8Format::E4M3>(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<torch::Tensor, torch::Tensor, torch::Tensor> 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<torch::Tensor> 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<torch::Tensor, torch::Tensor, torch::Tensor> 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<torch::Tensor, torch::Tensor, torch::Tensor> 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<FP8Format::E5M2, false, RowMajor, ColMajor>(
|
||||
p, stream.stream());
|
||||
@@ -288,9 +309,7 @@ std::tuple<torch::Tensor, torch::Tensor, torch::Tensor> linear_forward_fp8(
|
||||
|
||||
std::vector<int64_t> 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<torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor>
|
||||
@@ -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"),
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user