B200 / BF16 / MLPBlock / final implementation

SwiGLU MLP Drop-H 技术报告

最终方案很朴素:大 GEMM 继续交给 cuBLAS;forward 保存 preact,不保存 hidden activation h;SwiGLU forward 使用 autotuned tiled Triton kernel;backward 用 non-inplace Triton kernel 同时重算 h 和生成 grad_preact,并允许 autotune 搜索 tile。target-shape paired A/B 重测后,速度和 baseline 按 tie 处理,forward activation 显存结论是确定收益。

keep cuBLAS GEMM no weight pack cuBLAS linear1/linear2 drop saved h autotuned tiled SwiGLU forward non-inplace autotuned bwd bwd peak not targeted

目标不是改变 SwiGLU 数学,而是改变 autograd 保存哪些中间量。普通路径在 forward 后保存 preact 和 h;最终路径只保存 preact,让 h 在 linear2 之后释放。backward 需要 h 时,从 preact 重算,并额外分配 h 和 grad_preact,这样 Triton autotune 的每个候选都看到同一份原始输入。

after-forward memory
686.1 MiB
baseline 990.6 MiB,节省 304.5 MiB
peak memory
1905.1 MiB
non-inplace backward 为速度让路;bwd peak 不作为目标
full fwd+bwd
tie
paired median delta -1.84%,低于 3% 判定阈值
correctness
exact
final drop_h 的 out / grad_x / grad_w1 / grad_w2 max relative error 都为 0.0

SwiGLU MLP 的原始数学

MLPBlock 可以拆成三段:第一层 linear 得到两倍 hidden 宽度的 preactivation,SwiGLU 把它压回 hidden,第二层 linear 回到输出维度。

Forward 公式

preact = x @ W1.T
left, gate = split(preact, dim=-1)
h = left * SiLU(gate)
y = h @ W2.T

Backward 公式

grad_h = grad_y @ W2

dleft = grad_h * SiLU(gate)
dgate = grad_h * left * SiLU'(gate)
grad_preact = concat(dleft, dgate)

grad_x  = grad_preact @ W1
grad_W1 = grad_preact.T @ x
grad_W2 = grad_y.T @ h
SwiGLU MLP math diagram
普通 MLPBlock:h 是 linear2 的输入,baseline 会把它留给 backward。

Forward:保存 preact,不保存 h

forward 仍然用 PyTorch/cuBLAS 做两个大 GEMM,避免自写 GEMM 损失吞吐。SwiGLU activation 默认走 autotuned tiled Triton kernel,BM/BN/num_warps 由 Triton autotune 选择。custom autograd 保存 x, W1, W2, preact,但不保存 h;h 仍然会临时写回 HBM,作为 linear2 的输入,用完即可释放。

final forward drop h diagram
forward 的 memory win 来自 saved tensor set;tiled SwiGLU 只优化 activation kernel,不替换大 GEMM。

实现骨架

def forward(ctx, x, w1, w2):
    x2 = x.reshape(-1, x.shape[-1])

    preact = F.linear(x2, w1)  # cuBLAS
    h = tiled_swiglu(preact)   # temporary; not saved
    y = F.linear(h, w2)        # cuBLAS

    ctx.save_for_backward(x2, w1, w2, preact)
    return y.reshape(*x.shape[:-1], y.shape[-1])

张量生命周期

tensornormal baselinefinal drop_h
xsavedsaved
W1savedsaved
W2savedsaved
preact [M,2H]savedsaved
h [M,H]savedfreed

Backward:融合 h 重算和 grad_preact

backward 的关键路径是:先用 cuBLAS 算 grad_h = grad_y @ W2;然后 Triton kernel 一次性从 preact 和 grad_h 生成 h 与 grad_preact。最终版保持 non-inplace 语义:不覆盖 saved tensor,也不覆盖 grad_h,因此会额外分配 h 和 grad_preact。

Backward 实现骨架

def backward(ctx, grad_y):
    x, w1, w2, preact = ctx.saved_tensors
    grad_y = grad_y.reshape(-1, grad_y.shape[-1])

    grad_h = grad_y @ w2
    h, grad_preact = fused_h_and_grad_preact(
        preact,
        grad_h,
    )

    grad_w2 = grad_y.T @ h
    grad_x  = grad_preact @ w1
    grad_w1 = grad_preact.T @ x
    return grad_x, grad_w1, grad_w2

Fused kernel 伪代码

for tile in tiles(M, H):
    gh = load(grad_h)
    left = load(preact_left)
    gate = load(preact_gate)

    sig = sigmoid(gate)
    silu = gate * sig
    silu_prime = sig + silu * (1 - sig)

    h = bf16(left) * bf16(silu)
    dleft = gh * silu
    dgate = gh * left * silu_prime

    store(h_out, h)
    store(grad_preact_left, dleft)
    store(grad_preact_gate, dgate)
语义注意:当前最终版不原地覆盖 saved preact,也不原地覆盖 grad_h。这样代码语义更干净,也让 autotune 安全;代价是 backward peak memory 更高。

性能和显存对比

数据来自 B200 目标 shape [M,D,H,Dout] = [11136,3584,14336,3584]。重测使用同一张 B200、同一进程,A/B 随机交错运行,每个指标 40 个 paired samples。速度用 median 和 paired delta 判断;低于 3% 一律判为 tie。

最终判断:最终路径固定使用 non-inplace autotuned backward。它保留 forward 不保存 h 的显存收益,速度和 baseline 按 tie 处理;backward peak memory 会升高,但当前优化目标不再把 bwd peak 作为硬约束。

单个 MLP 目标 shape

path correctness fwd ms bwd ms full ms paired verdict after-fwd MiB peak MiB
normal baseline reference 2.2742 5.0725 7.4740 reference 990.6 1698.6
final drop_h + autotuned fwd/bwd exact 2.2332 5.0033 7.3664 vs normal tie: fwd -2.02%, bwd -1.30%, full -1.84% 686.1 1905.1

最终 paired A/B

comparison paired delta 结论
final drop_h vs normal baseline fwd -2.02%, bwd -1.30%, full -1.84% 低于 3% 阈值,端到端仍按 tie;peak memory 约 1905.1 MiB。

4-layer ParMLP stack memory

path fwd ms bwd ms full ms after-fwd MiB peak MiB
normal baseline 9.7416 21.4946 32.4014 3962.5 4670.5
final drop_h 9.7090 21.5666 32.0910 2744.5 3963.5

这张 stack 表用于看显存缩放,不作为速度胜负判据;速度胜负以单个 MLP paired A/B 表为准。

Exact 的口径:这里不是说数学上证明所有输入都 bitwise 恒等,而是 target-shape correctness 对比 baseline 时,final path 的 out、grad_x、grad_w1、grad_w2 最大相对误差都观测为 0.0。

取舍结论

这轮优化最后留下来的规律很清晰:大 GEMM 不要替换,小 activation/recompute kernel 可以自己写;省显存的核心是改变 saved tensor 生命周期,而不是把所有中间张量都消灭掉。

要保留

项目原因
cuBLAS GEMMlinear1、linear2、grad_x、grad_W1、grad_W2 都是主耗时,cuBLAS 吞吐最稳。
保存 preactbackward 精确需要 left/gate;不保存就要重算一次大 x @ W1.T。
不保存 h省掉 forward 后长期持有的 [M,H] activation,目标 shape 约 304.5 MiB。
autotuned tiled SwiGLU forward现在作为唯一最终路径;BM/BN/num_warps 由 autotune 选择,forward 单项在 target shape 上约快 2.02%。
fused activation backward一次 kernel 同时重算 h 并生成 grad_preact,减少 activation backward 的 kernel 数和中间调度开销。
non-inplace autotuned backward现在作为唯一最终路径;不覆盖 saved tensor,允许 autotune 安全搜索 tile。

不要作为默认

项目原因
packed weight会改变参数 layout 和集成 contract;当前 no-pack 已经拿到主要显存收益。
no-save 重算 linear1省掉 preact 但多一次大 GEMM,backward 速度不划算。
替换大 GEMM 的 fusionlinear1/linear2 或 backward GEMM 一旦离开 cuBLAS,吞吐损失会吃掉 activation fusion 的收益。

最终推荐

生产默认

使用 final drop_h:forward 省掉 saved h,SwiGLU forward 走 autotuned tiled Triton,backward 走 non-inplace autotuned fused h + grad_preact。速度对 baseline 按 tie 处理,forward activation 显存是主要收益。

回退路径

fallback 走普通 MLPBlock:继续保存 h,速度和数值语义最保守,只是不拿 forward activation 显存收益。