perf: transpose-quantize backward operands to route all gemms nt
- quantize gains out_layout (0 row-major / 1 transposed / 2 single-read dual-write); modes 1/2 run a new 32x32 smem-tile transpose kernel - backward feeds g8/w8T and g8T/x8T to trans_b=True gemms, dropping the NN-swap and TT crosswise kernels from training; fp8 weights keep the swap fallback - a 64x64 tile variant tied on the real step mix and was reverted; noted in the kernel header Benchmark: NVIDIA L20, 1.2B model, full train step fwd+bwd+CE - M=8192: fp8 551.8 -> 532.2 ms, 1.21x -> 1.26x vs bf16; M=2048 0.90x -> 0.95x - kernel-level grad_x +3.7..12.4%, grad_w +13.8..20.8%; layouts byte-exact, fp8 tests 36/36
This commit is contained in:
+15
-5
@@ -407,11 +407,21 @@ class _LinearFp8(torch.autograd.Function):
|
||||
meta.g.seed(g2, fmt)
|
||||
sg = meta.g.scale.clone()
|
||||
sw, sx = _sw_fwd, _sx_fwd
|
||||
g8, amax_g = quantize(g2, sg.reciprocal(), fmt)
|
||||
x8, _ = quantize(x.reshape(-1, x.size(-1)), sx.reciprocal(), fmt)
|
||||
w8 = w if _is_fp8(w.dtype) else quantize(w, sw.reciprocal(), fmt)[0]
|
||||
grad_x = mm_fp8(g8, w8, sg * sw).reshape(x.shape) # g8[m,n] @ w8[n,k]
|
||||
grad_w = mm_fp8(g8, x8, sg * sx, trans_a=True) # g8.T @ x8
|
||||
# Backward GEMMs route through the NT fast path via transposed
|
||||
# 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
|
||||
# crosswise kernel in the training path. g is consumed in both
|
||||
# orientations, so one dual-layout pass feeds both.
|
||||
g8, g8T, amax_g = quantize(g2, sg.reciprocal(), fmt, layout=2)
|
||||
x8T, _ = quantize(x.reshape(-1, x.size(-1)), sx.reciprocal(), fmt, layout=1)
|
||||
if _is_fp8(w.dtype):
|
||||
# Pre-quantized weight has no transposed copy: keep the swap
|
||||
# path for grad_x (grad_w is unaffected).
|
||||
grad_x = mm_fp8(g8, w, sg * sw).reshape(x.shape)
|
||||
else:
|
||||
w8T, _ = quantize(w, sw.reciprocal(), fmt, layout=1)
|
||||
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
|
||||
# bias-free linears must not pay the column-sum
|
||||
# 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
|
||||
|
||||
Reference in New Issue
Block a user