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:
2026-08-24 19:25:23 +08:00
parent 29e5f571af
commit 7da1439c9e
5 changed files with 114 additions and 44 deletions
+14 -7
View File
@@ -138,18 +138,25 @@ def mm_fp8(
return fp8_gemm(a, b, sa, sb, int(out_dtype == "e4m3"), out_scale) 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. """Pure FP8 linear forward: quantize x/w to ``fmt``, pre-quantized GEMM.
Returns ``(out, amax_x, amax_w)``. ``bias`` may be ``None``. Both Returns ``(out, amax_x, amax_w)``. ``bias`` may be ``None``. For static
operands share the same FP8 format (E4M3 by default; E5M2 for a fp8 inference, ``w`` and ``bias`` may arrive pre-quantized to ``fmt``
range-first configuration). (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): fmt8 = _fmt_dtype(fmt)
raise TypeError(f"fp8 forward requires bf16 inputs, got {x.dtype}/{w.dtype}") 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: if bias is None:
bias = torch.empty(0, device=x.device, dtype=x.dtype) 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"): def linear_backward_fp8(g, x, w, masks, sg, sw, sx, fmt: str = "e5m2"):
+2 -2
View File
@@ -75,16 +75,16 @@ struct FP8Params {
// the pre-quantized path. Scales are quantization steps (device scalars). // the pre-quantized path. Scales are quantization steps (device scalars).
const void* __restrict__ a_ptr = nullptr; const void* __restrict__ a_ptr = nullptr;
const void* __restrict__ b_ptr = nullptr; const void* __restrict__ b_ptr = nullptr;
const void* __restrict__ bias = nullptr;
const float* __restrict__ scale_a = nullptr; const float* __restrict__ scale_a = nullptr;
const float* __restrict__ scale_b = 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 // Output: BF16 or FP8 (E4M3). out_scale is the output quantization step
// (FP8 output only). // (FP8 output only).
void* __restrict__ out_ptr = nullptr; void* __restrict__ out_ptr = nullptr;
const float* __restrict__ out_scale = nullptr; const float* __restrict__ out_scale = nullptr;
// Fused forward extras: bias (may be null) and amax slots (may be null). // 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_a = nullptr;
float* __restrict__ amax_b = nullptr; float* __restrict__ amax_b = nullptr;
+21 -10
View File
@@ -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 // 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). // launcher dispatches to it there (see launch_fp8_gemm).
template <typename Traits, bool OutFp8 = false, typename LayoutA = RowMajor, typename LayoutB = RowMajor> template <typename Traits, bool OutFp8 = false, typename LayoutA = RowMajor, typename LayoutB = RowMajor>
__global__ void __launch_bounds__( __global__ void __launch_bounds__((Traits::kBlockM / 64) * (Traits::kBlockN / 32) * 32, 2)
(Traits::kBlockM / 64) * (Traits::kBlockN / 32) * 32, 2)
fp8_gemm_kernel(FP8Params p) { fp8_gemm_kernel(FP8Params p) {
using T8 = std::conditional_t<Traits::kIsE5M2, __nv_fp8_e5m2, __nv_fp8_e4m3>; using T8 = std::conditional_t<Traits::kIsE5M2, __nv_fp8_e5m2, __nv_fp8_e4m3>;
constexpr int kBlockM = Traits::kBlockM; constexpr int kBlockM = Traits::kBlockM;
@@ -453,36 +452,48 @@ fp8_gemm_kernel(FP8Params p) {
} }
const float output_scale = sa * sb; 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 #pragma unroll
for (int nt = 0; nt < 4; ++nt) { for (int nt = 0; nt < 4; ++nt) {
const int64_t col = output_col + nt * 8; 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 // Per-row store: FP8 packs two adjacent columns into one 16-bit
// write, BF16 into one 32-bit __nv_bfloat162 (single cvt+pack // write, BF16 into one 32-bit __nv_bfloat162 (single cvt+pack
// instruction); boundary or unaligned columns fall back to scalar // instruction); boundary or unaligned columns fall back to scalar
// converts so a pack never crosses the row edge or misaligns. // converts so a pack never crosses the row edge or misaligns.
auto store_out = [&](int64_t row, float v0, float v1) { auto store_out = [&](int64_t row, float v0, float v1) {
if (row >= m) return; if (row >= m) return;
const float r0 = v0 * output_scale + b0;
const float r1 = v1 * output_scale + b1;
if constexpr (OutFp8) { if constexpr (OutFp8) {
if (col + 1 < n) { if (col + 1 < n) {
*reinterpret_cast<unsigned short*>(out_fp8 + row * n + col) = *reinterpret_cast<unsigned short*>(out_fp8 + row * n + col) =
static_cast<unsigned short>(__nv_cvt_float2_to_fp8x2( 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)); __NV_SATFINITE, __NV_E4M3));
} else { } 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 { } else {
auto* dst = out_bf16 + row * n + col; auto* dst = out_bf16 + row * n + col;
if (col + 1 < n && if (col + 1 < n &&
(reinterpret_cast<uintptr_t>(dst) & 3) == 0) { (reinterpret_cast<uintptr_t>(dst) & 3) == 0) {
*reinterpret_cast<__nv_bfloat162*>(dst) = *reinterpret_cast<__nv_bfloat162*>(dst) =
__floats2bfloat162_rn(v0 * output_scale, __floats2bfloat162_rn(r0, r1);
v1 * output_scale);
} else { } else {
dst[0] = __float2bfloat16(v0 * output_scale); dst[0] = __float2bfloat16(r0);
if (col + 1 < n) if (col + 1 < n) dst[1] = __float2bfloat16(r1);
dst[1] = __float2bfloat16(v1 * output_scale);
} }
} }
}; };
+47 -25
View File
@@ -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, 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& 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) { int64_t k, int64_t a_ld, int64_t b_ld) {
p.a_ptr = a; p.a_ptr = a;
p.b_ptr = b; 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_a = sa.data_ptr<float>();
p.scale_b = sb.data_ptr<float>(); p.scale_b = sb.data_ptr<float>();
p.out_scale = out_scale ? out_scale->data_ptr<float>() : nullptr; 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_a = nullptr;
p.amax_b = nullptr; p.amax_b = nullptr;
p.m = static_cast<int>(m); 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)); : a_c.options().dtype(torch::kBFloat16));
FP8Params p; FP8Params p;
pack_gemm_params(p, a_c.data_ptr(), b_c.data_ptr(), out.data_ptr(), sa, sb, 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) if (a.scalar_type() == torch::kFloat8_e4m3fn)
dispatch_gemm<FP8Format::E4M3>(p, stream.stream(), out_fp8, ta, tb); dispatch_gemm<FP8Format::E4M3>(p, stream.stream(), out_fp8, ta, tb);
else 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( 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 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 // Pure FP8 forward: quantize x/w (fmt: 0 = E4M3, 1 = E5M2), then the
// pre-quantized GEMM; the dequantized BF16 output gets the bias added. // 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.is_cuda() && w.is_cuda(), "CUDA tensors required");
TORCH_CHECK(x.scalar_type() == torch::kBFloat16 && const auto f8opt = fmt ? torch::kFloat8_e5m2 : torch::kFloat8_e4m3fn;
w.scalar_type() == torch::kBFloat16, const bool w_prequant = w.scalar_type() == f8opt;
"x and w must be bf16"); 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"); TORCH_CHECK(x.device() == w.device(), "x and w must be on the same device");
check_scale(sx, x, "sx"); check_scale(sx, x, "sx");
check_scale(sw, x, "sw"); 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); 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"); 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 has_bias = bias.defined() && bias.numel() > 0;
const bool b_prequant = has_bias && bias.scalar_type() == f8opt;
if (has_bias) { if (has_bias) {
TORCH_CHECK(bias.is_cuda() && bias.device() == x.device() && TORCH_CHECK(bias.is_cuda() && bias.device() == x.device() &&
bias.scalar_type() == torch::kBFloat16 && bias.numel() == n &&
bias.numel() == n, (bias.scalar_type() == torch::kBFloat16 || b_prequant),
"bias must be CUDA bf16 with shape [N]"); "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 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_x = torch::zeros({1}, x.options().dtype(torch::kFloat32));
auto amax_w = 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()); 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(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; FP8Params p;
// Forward is the NT layout: A = x8 [M,K] (a_ld = k), B = w8 [N,K] // 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, 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) { if (fmt) {
launch_fp8_gemm<FP8Format::E5M2, false, RowMajor, ColMajor>( launch_fp8_gemm<FP8Format::E5M2, false, RowMajor, ColMajor>(
p, stream.stream()); 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); std::vector<int64_t> shape(x.sizes().begin(), x.sizes().end() - 1);
shape.push_back(n); shape.push_back(n);
auto out_r = out.reshape(shape); return {out.reshape(shape), amax_x, amax_w};
if (has_bias) out_r = out_r + bias;
return {out_r, amax_x, amax_w};
} }
std::tuple<torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor> 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}); auto grad_input_2d = grad_input.reshape({m, k});
FP8Params gp; FP8Params gp;
pack_gemm_params(gp, g8.data_ptr(), w8.data_ptr(), pack_gemm_params(gp, g8.data_ptr(), w8.data_ptr(),
grad_input_2d.data_ptr(), sg, sw, nullptr, m, k, n, n, grad_input_2d.data_ptr(), sg, sw, nullptr, nullptr,
k); nullptr, m, k, n, n, k);
run_bwd_gemm(gp, false, false); 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 = // 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); quantize(x_c, x8, sx, nullptr);
FP8Params gp; FP8Params gp;
pack_gemm_params(gp, g8.data_ptr(), x8.data_ptr(), 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); run_bwd_gemm(gp, true, false);
} }
if (!masks[0] && !masks[1]) { if (!masks[0] && !masks[1]) {
@@ -402,9 +422,11 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
"the operand layout (default 0/0 = a@b)"); "the operand layout (default 0/0 = a@b)");
m.def("linear_forward_fp8", &linear_forward_fp8, py::arg("x"), 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("w"), py::arg("bias"), py::arg("sx"), py::arg("sw"),
py::arg("fmt") = 0, py::arg("fmt") = 0, py::arg("bias_scale") = py::none(),
"Pure FP8 linear forward: quantize x/w, pre-quantized GEMM; " "Pure FP8 linear forward: quantize x/w, pre-quantized GEMM with the "
"returns (out, amax_x, amax_w)"); "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"), 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("x"), py::arg("w"), py::arg("masks"), py::arg("sg"),
py::arg("sw"), py::arg("sx"), py::arg("fmt"), py::arg("sw"), py::arg("sx"), py::arg("fmt"),
+30
View File
@@ -146,6 +146,36 @@ def test_linear_backward_e5m2_gradients():
torch.testing.assert_close(amax_g, grad.abs().amax().float().reshape(1)) 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 @skip_no_fp8
def test_fp8_linear_backward_outside_autocast(): def test_fp8_linear_backward_outside_autocast():
"""aten::linear records an fp8 autograd node inside fp8_autocast; the """aten::linear records an fp8 autograd node inside fp8_autocast; the