跳到正文

谁偷走了 5090 的算力和显存:一步 Transformer 训练的 roofline 侦查

我在一张 RTX 5090 上把一步 Transformer 训练拆开,分别量了时间和显存。结果和预想的不太一样:拖慢速度的不是矩阵乘,占显存最多的也不是权重。两边查到最后,都落在 attention 里的两个 seq × seq 矩阵上,一个是分数矩阵 S = QKᵀ(每个 query 对每个 key 的打分),一个是 S 按行过 softmax 之后的注意力权重 P,最后输出是 PV。

文章分四步走:第 1 节准备工具 roofline;第 2 节算一步训练的总账,找出时间和显存的大头;第 3 节逐个 op 拆开,看每一步算了多少、读写了多少显存、为反向存了什么,先拿最简单的 RMSNorm 走一遍,再用同样的方法拆 attention;第 4 节试 bf16 和 activation checkpoint 能省多少。怎么把 S、P 彻底去掉,留给下一篇 FlashAttention 1–4。GPU 的硬件和 Triton 的基本写法放在系列的上一篇 GPU 与 Triton 入门。

模型是我自己写的 Transformer LM(RMSNorm、RoPE、SwiGLU,pre-norm),一共五档:small 0.13B、medium 0.42B、large 0.97B、xl 3.41B,以及 10B(实际 12.83B 参数)。

环境:RTX 5090 32 GB,torch 2.11.0+cu130。fp32 基准关掉 tf32(allow_tf32=False)。除注明外 batch 4、seq 512,预热 5 步、计时 10 步。「full step」指前向、反向加 optimizer 的一整步,「前向 + 反向」不含 optimizer。显存一律 GiB = 2³⁰ B(max_memory_allocated() / 1024³)。


1 Roofline:一把尺子
#

一个 op 跑多快,取决于它要算多少和要搬多少。

算的量用 FLOPs(浮点运算次数)数。矩阵乘 [M, K] × [K, N] 要 2·M·N·K 次(M·N 个输出,每个做 K 次乘加);逐元素 op(加、乘、exp、mask)每个元素只算一到几十次,比同样大小的矩阵乘少几个数量级。GPU 每秒最多能做的浮点运算次数叫峰值算力,记作 $\pi$,单位 FLOPS。5090 的 fp32 峰值是 $\pi$ = 1.05e14 FLOPS,用 bf16 Tensor core 时翻倍到 2.1e14。

搬的量是 op 读写显存的字节数:输入要从显存读进来,结果要写回去。显存每秒最多能读写的字节数叫带宽,记作 $\beta$,5090 是 $\beta$ = 1.79e12 B/s。一个 op 的耗时不会低于算的时间和搬的时间中较大的那个:

$t \ge \max\left(\dfrac{\mathrm{FLOPs}}{\pi},\ \dfrac{\mathrm{bytes}}{\beta}\right)$

算和搬的比值叫算术强度 $I = \mathrm{FLOPs} / \mathrm{bytes}$,即每搬 1 字节做多少次运算。把上式改写成「最多能跑到多少 FLOPS」,就是 roofline:

$\mathrm{FLOPS}_{\max}(I) = \min(\pi,\ I \cdot \beta)$

在 log-log 坐标上它像一个屋顶(图 1-1):左边是斜坡,被带宽限制;右边是平顶,被算力限制;拐角 $I^* = \pi / \beta$ 叫 ridge point,5090 fp32 是 58 FLOPs/B。落在拐角左边的 op 是 memory-bound,耗时由搬数据决定,减少 FLOPs 没有用;落在右边的是 compute-bound,耗时由算力决定。衡量 op 跑得好不好也分两种:compute-bound 的看 MFU(实际 FLOPS / $\pi$),memory-bound 的看 MBU(实际带宽 / $\beta$)。

0.11101001000100001e111e121e131e141e15算术强度 I(FLOPs/B,对数轴)可达算力(FLOPS,对数轴)拐点旁的数字是 ridge point I*(FLOPs/B)fp32:ridge point = 1.05e14 / 1.792e12 = 58 FLOPs/B58fp32 π = 1.05e14bf16:ridge point = 2.1e14 / 1.792e12 = 117 FLOPs/B117bf16 2.1e14fp8:ridge point = 4.19e14 / 1.792e12 = 234 FLOPs/B234fp8 4.19e14nvfp4:ridge point = 1.68e15 / 1.792e12 = 935 FLOPs/B935nvfp4 1.68e15带宽 β = 1.79e12 B/smemory-boundcompute-bound
图 1-1 RTX 5090 各精度的 roofline:斜线是带宽,平线是峰值算力(dense,boost clock 2407 MHz,来自 NVIDIA RTX Blackwell 白皮书;Tensor core 按 fp32 累加)。

一个 op 在拐角哪一边,用张量形状就能估。逐元素 op 在 fp32 下每个元素算 1 次、读写 8 字节,$I \approx 0.13$,远在拐角左边。矩阵乘的 $I$ 由 M、N、K 里最小的那个决定,fp32 下不超过它的一半。Linear 的三个维度都上千(medium、seq 1024 时 FFN 第一层是 [4096, 1024] × [1024, 4096]),$I \approx 340$,在拐角右边;attention 里的 QKᵀ 和 PV 都有一个维度是 d_head(64),$I$ 只有 28,在拐角左边。所以 QKᵀ 和 PV 虽然是矩阵乘,在 attention 里也是 memory-bound。3.2 节会把实测的 op 放到 fp32 的 roofline 上看(图 3-5)。


2 总账:时间和显存花在哪
#

2.1 时间:矩阵乘没在偷懒
#

Transformer 的 FLOPs 几乎都在 Linear 上。一个 Linear 前向只做一次矩阵乘 $X_L = X_{L-1} W_L$;反向收到误差 $\nabla X_L$ 后要做两次:算参数梯度 $\nabla W_L = X_{L-1}^{\top} \nabla X_L$,再算传给上一层的 $\nabla X_{L-1} = \nabla X_L W_L^{\top}$,每次都和前向一样大(图 2-1)。

Forward Pass(前向)Backward Pass(反向)输入XL−1计算XL=XL−1·WL权重WL输出XL进入深层 Layer L+1,等误差传回接收误差∇XL计算参数梯度∇WL=XL−1T· ∇XL计算激活梯度∇XL−1= ∇XL·WLT进入浅层 Layer L−1XL−1留到反向(第 3 节的 A)XL−1留到反向(第 3 节的 A)WL是参数,本来就在显存里WL是参数,本来就在显存里
图 2-1 一个 Linear 的前向与反向:前向从左边往下,误差从右边传回;虚线是反向要从前向拿的东西。

每个参数对每个 token 做一次乘加,也就是 2 FLOPs,所以前向每个 token 约 2N FLOPs(N 是参数量),反向 4N,一步一共 6N × token 数,这里是 batch 4 × seq 512 = 2048 个 token。attention 的 QKᵀ 和 PV 没有参数,不在 6N 里,seq 512 时只占 2–4%,可以忽略。实测也是这样,反向耗时差不多是前向的两倍(图 2-2),因为反向要把 weight grad 和 activation grad 各算一遍。

前向反向optimizersmall0.13Bsmall 前向:17.2 ms(31%)31%small 反向:34.8 ms(62%)62%small optimizer:3.7 ms(7%)55.7 ms / 步MFU 27%medium0.42Bmedium 前向:51.1 ms(31%)31%medium 反向:103.3 ms(62%)62%medium optimizer:12.9 ms(8%)167.4 ms / 步MFU 31%large0.97Blarge 前向:118.1 ms(32%)32%large 反向:227.8 ms(61%)61%large optimizer:26.9 ms(7%)372.8 ms / 步MFU 31%
图 2-2 一步训练里前向、反向、optimizer 的耗时占比(fp32,batch 4,seq 512),右侧是每步耗时和 MFU。

按 6N 算,medium 和 large 的 MFU 只有 31%(实际 3.2e13 FLOPS,fp32 峰值 1.05e14),small 是 27%。矩阵乘本身不慢,单个大矩阵乘能跑到峰值的 64%;问题是所有矩阵乘加起来只占一步 GPU 时间的 60%,剩下 40% 花在几乎不做计算的逐元素 kernel 上,比如 norm、softmax、mask、激活函数。

2.2 显存:权重只是小头
#

fp32 + AdamW 训练时,每个参数要占 16 B:权重 4 B、梯度 4 B、Adam 的 m 和 v 各 4 B,此外还有前向为反向留下的 activation。xl 光这部分就要 50.8 GiB,5090 放不下;10B 在建模型时就 OOM 了。下面用这几个记号:

  • W 全部权重;G 全部梯度 .grad,大小等于 W;Adam 的 m、v 合计 2W,第一步之后常驻;
  • A 前向为反向存下的张量(saved tensors);T 当前层的临时量,算完即释放。

实测 full step 的峰值里只有权重、Adam 状态和 A,没有梯度(图 2-3)。

W 权重Adam m、vA activationT 临时量0102030GiBsmall:W 0.48 GiBsmall:Adam 0.96 GiBsmall:A 3.50 GiBA 3.5small:T 0.10 GiB5.04 GiBsmall 0.13Bmedium:W 1.58 GiBmedium:Adam 3.16 GiBAdam 3.2medium:A 8.91 GiBA 8.9medium:T 0.09 GiB13.74 GiBmedium 0.42Blarge:W 3.61 GiBW 3.6large:Adam 7.22 GiBAdam 7.2large:A 16.58 GiBA 16.6large:T 0.10 GiB27.51 GiBlarge 0.97B
图 2-3 full step 的实测峰值显存(batch 4,seq 512)。A = 带梯度的前向峰值 − W,顶上 ~0.1 GiB 是 T。

梯度不在峰值里,是因为 G 和 A 此消彼长。反向走完 j 层(共 L 层)时,显存里是:

$M(j) = W + G \cdot \dfrac{j}{L} + A \cdot \dfrac{L-j}{L} + T$

A 在前向一层层攒起来,反向再一层层释放;G 正好相反,前向时还不存在,.grad 要等反向算到那个参数才分配,optimizer step 结束后又被 zero_grad(set_to_none=True) 释放。一个涨一个降,M(j) 是一条直线,最高点只可能在两端。full step 里还有一直占着的 Adam 状态,m 和 v 各和权重一样大,共 2W,和权重加起来常驻 3W。所以 full step 的峰值是

$\mathrm{peak}_{\mathrm{full}} \approx 3W + \max(A,\ G)$

拿 large 验算:3 × 3.61 + 16.58 = 27.4 GiB,实测 27.51 GiB,差的 0.1 GiB 就是 T。正常训练 token 多,A 比 G 大,峰值出现在前向刚结束时;token 很少时才反过来,比如 xl 在 seq 128(图 2-4 左)。

WAGT×逐层实测xl,seq 128(G > A)0510152025WAGT××××××××××××××××××××××××××××××××××××××××××××××××××××××××××××××××峰值 25.56 GiB开始前向结束反向结束W = G = 12.7,A = 5.3 GiBsmall,seq 512(A > G)01234WAGT××××××××××××××××××××××××峰值 4.08 GiB开始前向结束反向结束W = G = 0.48,A = 3.41 GiBGiB
图 2-4 一步前向 + 反向的显存(fp32):色带按 M(j) 用实测的 W、A、G、T 堆叠,× 是逐层实测值。

总账算下来有两条线索:时间上,40% 花在逐元素 kernel 上;显存上,峰值主要由 A 决定,而且四项里只有 A 随 seq 变大。下面把单个 op 拆开来看。


3 拆解:逐个 op 看
#

这一节对每个 op 问三件事:算了多少 FLOPs,读写了多少字节显存,为反向存了哪些张量。前两件看张量形状就能算出来;第三件由 autograd 决定,要实测。

PyTorch 的 torch.autograd.graph.saved_tensors_hooks(pack, unpack) 就是用来看第三件事的。它是一个上下文管理器:在它里面跑前向时,autograd 每存下一个张量就调用一次 pack(t),保存的是 pack 的返回值;反向每用到一个存下的张量就调用一次 unpack,收到的就是当初 pack 的返回值。下面的两个函数都原样返回张量,只是打印出来,并用 data_ptr() 给每块内存编号,编号相同就是同一块内存:

blocks, count = {}, [0]  # data_ptr -> 块编号:编号相同就是同一块内存

def block(t):
    return blocks.setdefault(t.data_ptr(), "ABCDEFGHIJ"[len(blocks)])

def pack(t):  # 前向每存下一个张量调用一次,返回值会被保存
    count[0] += 1
    print(f"Saving  {count[0]}  {list(t.shape)}  {str(t.dtype)[6:]}  块 {block(t)}")
    return t

def unpack(t):  # 反向每取出一个张量调用一次,收到的就是 pack 的返回值
    print(f"Loading    块 {block(t)}")
    return t

def show(fn, *inputs):
    blocks.clear(); count[0] = 0
    with torch.autograd.graph.saved_tensors_hooks(pack, unpack):
        out = fn(*inputs)  # 前向:触发 pack
    out.sum().backward()   # 反向:触发 unpack

3.1 节拿最简单的 RMSNorm 把这套方法走一遍,3.2 节再用到 attention 上。

3.1 RMSNorm:融合省掉 x̂
#

RMSNorm 算的是 $y = w \odot (x \cdot r)$,其中 $r = (\tfrac{1}{d}\sum_j x_j^2 + \epsilon)^{-1/2}$ 每行一个数。eager 模式下它拆成 5 个 op(fp32,x: [4, 512, 2560],一份 x 是 20 MiB):

rms = torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + eps)  # ① pow ② mean ③ rsqrt
x_hat = x * rms                                            # ④
y = weight * x_hat                                         # ⑤

把上面三行包成函数 rmsnorm(x, w),用 show(rmsnorm, x, w) 跑一遍,打印如下。块 A 是 x,B 是 r,C 是 $\hat{x}$,D 是 w:

Saving  1  [4, 512, 2560]  float32  块 A
Saving  2  [4, 512, 1]  float32  块 B
Saving  3  [4, 512, 1]  float32  块 B
Saving  4  [4, 512, 2560]  float32  块 A
Saving  5  [4, 512, 2560]  float32  块 C
Saving  6  [2560]  float32  块 D
Loading    块 C
Loading    块 D
Loading    块 B
Loading    块 A
Loading    块 B
Loading    块 A

打印的结果说明了 autograd 存张量的规则:一个 op 的局部偏导里用到谁,前向就存谁;偏导是常数就什么都不存。存的是引用,不是拷贝,所以本来就在显存里的输入 x 和参数 w 不额外占显存。6 次 Saving 只落在 4 块内存上,新占显存的只有 $r$(8 KiB)和 $\hat{x}$(20 MiB),见表 3-1 和图 3-1。

表 3-1 RMSNorm 五个 op 为反向存的张量(eager)
op反向要的偏导存下新占显存
① $x^2$$\partial x^2/\partial x = 2x$$x$0
② $v=\tfrac1d\sum x^2$$\partial v/\partial x^2 = \tfrac1d$,常数不存0
③ $r=(v+\epsilon)^{-1/2}$$\partial r/\partial v = -\tfrac12 r^3$$r$8 KiB
④ $\hat{x}=x\cdot r$$\partial\hat{x}/\partial x = r$,$\partial\hat{x}/\partial r = x$$r$、$x$0
⑤ $y=w\odot\hat{x}$$\partial y/\partial w = \hat{x}$,$\partial y/\partial\hat{x} = w$$\hat{x}$、$w$20 MiB

④ 存的 $r$、$x$ 和 ③、① 存的是同一块内存,所以不再占显存。新占显存按 fp32 算:$r$ 是 [4, 512, 1],$\hat{x}$ 是 [4, 512, 2560]。

① x²1 FLOP/元素② mean1 FLOP/元素③ rsqrt每行 2 FLOPs④ x · r1 FLOP/元素⑤ w ⊙ x̂1 FLOP/元素x …904020 MiB,输入x²20 MiB,临时v8 KiB,临时r …6b00+8 KiBx̂ …3c00+20 MiBw …1000参数① 反向② 反向③ 反向④ 反向⑤ 反向ydydx前向箭头指向 op 是读,指向显存是写前向箭头指向 op 是读,指向显存是写显存显存反向反向
图 3-1 RMSNorm(eager,x 是 20 MiB):中间一排是显存里的张量,框宽按实际字节数线性画(KiB 级的只剩一条细线),每块只画一次,同色是同一块内存(ptr 相同);粗实线框新占显存,细实线框本来就在、只被引用(x 是上一层的输出,w 是参数,都会一直留着),灰色虚线框是用完即释放的临时量。前向的实线箭头指向 op 是读、指向显存是写;反向沿虚线箭头读回存下的张量。

时间上,RMSNorm 每个元素只算约 4 次,5 个 op 却要读写约 140 MiB 显存:①、④、⑤ 各读 20 MiB、写 20 MiB,② 读 20 MiB。$I \approx 0.14$,是 memory-bound。显存上,它为反向多存了一份 $\hat{x}$。

$\hat{x}$ 其实不用存:它就是 $x \cdot r$,反向时用 $x$ 和 $r$ 重算一次就有。eager 模式做不到,因为每个 op 单独执行,⑤ 的反向只知道自己需要 $\hat{x}$,不知道它能由 $x$ 和 $r$ 算出来。用 torch.compile 把 ①–⑤ 编译成一个前向 kernel、反向编成 3 个 kernel 后,只存 $x$、$w$、$r$:

Saving  1  [4,512,2560]  grad_fn=None  ptr=…3c80   # x
Saving  2  [2560]        grad_fn=None  ptr=…b3c0   # w
Saving  3  [4,512,1]     grad_fn=None  ptr=…f9c0   # r
Loading    1 → 2 → 3(与 Saving 同序)

反向要算 $\nabla x = r\,\big(g - \hat{x} \cdot \mathrm{mean}(g \odot \hat{x})\big)$,其中 $g = w \odot \nabla y$。式子里的 $\hat{x}$ 都在反向 kernel 里用 $x \cdot r$ 现算(图 3-2),前后对比见表 3-2。

fused forward:①–⑤ 一个 kernel~4 FLOPs/元素x …3c8020 MiB,输入r …f9c0+8 KiBw …b3c0参数反向 3 个 kernel(dx 1 个,dw 2 个),x̂ = x · r 现场重算xydydx前向箭头指向 op 是读,指向显存是写前向箭头指向 op 是读,指向显存是写显存显存反向反向
图 3-2 融合后的 RMSNorm(画法同图 3-1):前向只读 x、w,写 r 和 y,x²、v、x̂ 都不进显存;反向读回 x、w、r,现场重算 x̂。
表 3-2 RMSNorm:eager 与 torch.compile 融合
eager融合后
前向 kernel6 个(③ 的 +ε 和 rsqrt 各一个)1 个
反向 kernel13 个3 个(dx 1 个,dw 两段归约 2 个)
为反向存的张量$x$、$w$、$r$、$\hat{x}$$x$、$w$、$r$
新占显存$r + \hat{x}$ ≈ 20 MiB$r$ ≈ 8 KiB
前向读写显存~140 MiB~40 MiB(只读 $x$、写 $y$)
FLOPs / 元素前向 ~4前向 ~4,反向多 1(重算 $\hat{x} = x\cdot r$)
前向算术强度 $I$~4 / 28 B ≈ 0.14~4 / 8 B ≈ 0.5

kernel 数和存的张量是实测(torch.profiler,不含 memset、拷贝和 .grad 累加),FLOPs 和读写是纸面计数。I 按每个元素算:eager 每元素读写 7 次 × 4 B = 28 B(共 ~140 MiB),融合后只读 x、写 y,8 B。两者都远低于 ridge point 58,仍是 memory-bound。

融合没有让 RMSNorm 变成 compute-bound。把两种写法放到 fp32 roofline 上(图 3-3),I 从 0.14 移到 0.5,两个点都贴着带宽斜线,MBU 都在 80% 左右;要读写的字节少了 3.5 倍,前向耗时也从 816 µs 降到 236 µs,正好快 3.5 倍。

0.11101001000100001e111e121e131e14算术强度 I(FLOPs/B,对数轴)可达算力(FLOPS,对数轴)fp32:ridge point = 1.05e14 / 1.792e12 = 58 FLOPs/Bridge point I* = 58峰值 π = 1.05e14 FLOPS带宽 β = 1.79e12 B/smemory-boundcompute-boundRMSNorm eager(5 个 op):I = 0.143 FLOPs/B,实测 2.06e11 FLOPS,MBU 80%eagerRMSNorm torch.compile(1 个 kernel):I = 0.5 FLOPs/B,实测 7.1e11 FLOPS,MBU 79%融合后
图 3-3 RMSNorm 融合前后在 fp32 roofline 上的位置(前向,x: [32, 512, 2560],160 MiB,比 L2 大,避免数据留在缓存里)。读写字节按 eager 每元素 28 B、融合后 8 B 算,悬停可看数值。

融合版也可以手写成 Triton(下面折叠的代码,eager 版就是本节开头那三行;Triton 的基本写法见上一篇 GPU 与 Triton 入门),思路和 torch.compile 一样:前向一个 kernel,一个 program 算一行,①–⑤ 都在 register 里做完,只写出 y 和 r。反向比 torch.compile 少两个 kernel:dx 是按行归约,dw 却要把所有行加起来,torch.compile 为 dw 单独拆了两段归约;手写版让每个 program 读回 x、w、r,现场重算 x̂,算完自己这一行的 dx,再用 atomic_add 把这一行对 dw 的贡献直接累加上去,反向就只有一个 kernel。

Triton 融合:前向 1 个 kernel,反向 1 个 kernel
import torch
import triton
import triton.language as tl

@triton.jit
def rmsnorm_fwd(X, W, Y, R, D, eps, BLOCK: tl.constexpr):
    row = tl.program_id(0)                       # 一个 program 算一行
    cols = tl.arange(0, BLOCK)
    mask = cols < D
    x = tl.load(X + row * D + cols, mask=mask, other=0.0)
    w = tl.load(W + cols, mask=mask, other=0.0)
    r = tl.rsqrt(tl.sum(x * x, axis=0) / D + eps)  # ①②③ 在 register 里做完
    tl.store(R + row, r)                           # 只存 r:每行 4 B
    tl.store(Y + row * D + cols, w * (x * r), mask=mask)  # ④⑤,x̂ 不落显存

@triton.jit
def rmsnorm_bwd(DY, X, W, R, DX, DW, M, D, ROWS: tl.constexpr, BLOCK: tl.constexpr):
    cols = tl.arange(0, BLOCK)
    mask = cols < D
    w = tl.load(W + cols, mask=mask, other=0.0)
    dw = tl.zeros([BLOCK], dtype=tl.float32)
    for i in range(ROWS):                        # 一个 program 算 ROWS 行
        row = tl.program_id(0) * ROWS + i
        m = mask & (row < M)
        x = tl.load(X + row * D + cols, mask=m, other=0.0)
        dy = tl.load(DY + row * D + cols, mask=m, other=0.0)
        r = tl.load(R + row, mask=row < M, other=0.0)
        x_hat = x * r                            # 现场重算 x̂
        g = w * dy
        dx = r * (g - x_hat * tl.sum(g * x_hat, axis=0) / D)
        tl.store(DX + row * D + cols, dx, mask=m)
        dw += dy * x_hat
    tl.atomic_add(DW + cols, dw, mask=mask)      # 各 program 的 dw 部分和累加到一起

class RMSNorm(torch.autograd.Function):
    @staticmethod
    def forward(ctx, x, weight, eps=1e-5):
        D = x.shape[-1]
        x2 = x.reshape(-1, D)
        M = x2.shape[0]
        y = torch.empty_like(x2)
        r = torch.empty(M, device=x.device, dtype=torch.float32)
        rmsnorm_fwd[(M,)](x2, weight, y, r, D, eps, BLOCK=triton.next_power_of_2(D))
        ctx.save_for_backward(x2, weight, r)     # 存 x、w、r,不存 x̂
        return y.view_as(x)

    @staticmethod
    def backward(ctx, dy):
        x2, weight, r = ctx.saved_tensors
        M, D = x2.shape
        dx = torch.empty_like(x2)
        dw = torch.zeros(D, device=x2.device, dtype=torch.float32)
        ROWS = 16
        rmsnorm_bwd[(triton.cdiv(M, ROWS),)](dy.reshape(M, D).contiguous(), x2, weight, r, dx, dw, M, D,
                                             ROWS=ROWS, BLOCK=triton.next_power_of_2(D))
        return dx.view(dy.shape), dw, None

融合解决了两件事:几个 op 合成一个 kernel,中间结果不再进出显存;能重算的张量不存,反向时再算。3.2 节的 attention 和 4.2 节的 checkpoint 都会再遇到这两件事。

3.2 Attention:都在搬 S 和 P
#

一层 attention 的完整公式是

$O = \mathrm{softmax}\!\left(\dfrac{QK^{\top}}{\sqrt{d}} + M\right) V$

$Q$、$K$、$V$ 的形状都是 [b, h, seq, d](d 是 d_head),$M$ 是 causal mask(未来位置为 $-\infty$)。eager 模式下它拆成 5 步:① 算分数 $S = QK^{\top}$;② 除以 $\sqrt{d}$;③ 加 mask;④ 按行 softmax 得到 $P$;⑤ $O = PV$。S、P 的形状都是 [b, h, seq, seq],O 和 Q 一样大。

和 RMSNorm 一样,逐步看它算了多少、读写了多少显存、为反向存了什么(图 3-4,medium、seq 1024:b = 4,h = 16,d = 64)。一份 S 或 P 是 4 × 16 × 1024 × 1024 × 4 B = 256 MiB,而 Q、K、V、O 各只有 16 MiB。

eager 的写法如下,softmax 按公式拆成 5 个 kernel:

def attention(q, k, v, mask):
    s = q @ k.transpose(-2, -1) / math.sqrt(q.shape[-1])     # ① ②
    s = s.masked_fill(mask, float("-inf"))                    # ③
    m = s.max(dim=-1, keepdim=True).values                    # ④ softmax 的 5 个 kernel
    e = torch.exp(s - m)
    p = e / e.sum(dim=-1, keepdim=True)
    return p @ v                                              # ⑤

用 show 跑一遍,打印如下。块 A 是 K(转置后的 view),B 是 Q,C 是 mask,D 是每行 max 的下标,E 是 e = exp(S − m),F 是行和 Σ,G 是 V,H 是 P:

Saving  1  [64, 64, 1024]  float32  块 A
Saving  2  [64, 1024, 64]  float32  块 B
Saving  3  [1024, 1024]  bool  块 C
Saving  4  [4, 16, 1024, 1]  int64  块 D
Saving  5  [4, 16, 1024, 1024]  float32  块 E
Saving  6  [4, 16, 1024, 1]  float32  块 F
Saving  7  [4, 16, 1024, 1024]  float32  块 E
Saving  8  [64, 1024, 64]  float32  块 G
Saving  9  [64, 1024, 1024]  float32  块 H
Loading    块 G
Loading    块 H
Loading    块 F
Loading    块 E
Loading    块 E
Loading    块 D
Loading    块 C
Loading    块 A
Loading    块 B

按 RMSNorm 那条规则逐个 op 对一遍(表 3-3):9 次 Saving 里,Q、K、V、mask 本来就在显存里,只是引用;softmax 的 5 个 kernel 里,max 只把梯度传给最大值所在的位置,存下每行最大值的下标;exp 的导数就是它自己的输出,除法要用分子和分母,所以存下 e 和 Σ;⑤ 要用 P 和 V。新占显存的是 e 和 P 两个 seq × seq 张量,各 256 MiB,加起来是 Q、K、V、O 总和的 8 倍。

表 3-3 attention 各 op 为反向存的张量(eager,medium,seq 1024)
op反向要的偏导存下新占显存
① $S = QK^{\top}$$\partial S/\partial Q = K$,$\partial S/\partial K = Q$$Q$、$K$0
② $\div\sqrt{d}$$1/\sqrt{d}$,常数不存0
③ $+M$被 mask 的位置梯度为 0mask0
④ max只有最大值的位置有梯度下标0.5 MiB
④ $e = \exp(S - m)$$\partial e/\partial S = e$$e$256 MiB
④ $P = e / \Sigma$$\partial P/\partial e = 1/\Sigma$,$\partial P/\partial \Sigma = -e/\Sigma^2$$e$、$\Sigma$0.25 MiB
⑤ $O = PV$$\partial O/\partial P = V$,$\partial O/\partial V = P$$P$、$V$256 MiB

存下的张量是实测。④ 的减 max 和求和偏导是常数,不存;除法存的 e 和 exp 存的是同一块内存。

① QKᵀ② ÷√d③ +Mmax− mexp求和÷ Σ⑤ PVM1 MiBQ16 MiBK16 MiBS256 MiBS/√d256 MiBS+M256 MiB下标+0.5 MiBm0.25 MiBS−m256 MiBe+256 MiBΣ+0.25 MiBP+256 MiBV16 MiB① QKᵀ② ÷√d③ +Mmax− mexp求和÷ Σ⑤ PV④ softmax,5 个 kernelOdOdQ dK前向箭头指向 op 是读,指向显存是写前向箭头指向 op 是读,指向显存是写显存显存反向反向
图 3-4 eager attention 一层(medium,seq 1024,画法同图 3-1)。框下是每块的大小:S、S/√d、S+M、S−m、e、P 都是 [b, h, seq, seq],m 和 Σ 每行一个数。max 同时写出每行最大值的下标(int64),反向只用它,m 用完即释放。粗实线框是为反向新存下的,细实线框是本来就在、只被引用的 Q、K、V、mask,灰色虚线框是用完即释放的临时量。

时间上,FLOPs 集中在 ① 和 ⑤ 两个矩阵乘上,读写却集中在 ② ③ ④ 上:② 和 ③ 各把整个 256 MiB 的 S 读一遍、写一遍,④ 读写得更多。放到 roofline 上(图 3-5),attention 的 op 全在斜坡上,包括 QKᵀ 和 PV 这两个矩阵乘。它们的 MBU 已经有 57–86%,kernel 本身没多少优化空间,要更快只能少搬数据。

0.11101001000100001e111e121e131e14算术强度 I(FLOPs/B,对数轴)可达算力(FLOPS,对数轴)fp32:ridge point = 1.05e14 / 1.792e12 = 58 FLOPs/Bridge point I* = 58峰值 π = 1.05e14 FLOPS带宽 β = 1.79e12 B/smemory-boundcompute-boundLinear(FFN w1):I = 341 FLOPs/B,实测 6.87e13 FLOPS,MFU 64%LinearS = QKᵀ:I = 28.4 FLOPs/B,实测 2.86e13 FLOPS,MBU 57%QKᵀO = PV:I = 30.1 FLOPs/B,实测 3.73e13 FLOPS,MBU 70%PVsoftmax(5 个 kernel):I = 0.844 FLOPs/B,实测 1.27e12 FLOPS,MBU 84%softmaxS / √d:I = 0.125 FLOPs/B,实测 1.92e11 FLOPS,MBU 86%S / √d
图 3-5 RTX 5090 fp32 的 roofline,以及 medium、seq 1024 时一层 attention 里实测的 op(另放一个 Linear 作对照)。causal mask 没有 FLOPs,不在图上;悬停可看数值。

搬得最多的是 ④ softmax。它对 S 的每一行算 $P_{ij} = e^{S_{ij} - m_i} / \sum_k e^{S_{ik} - m_i}$,其中 $m_i = \max_k S_{ik}$,减掉行最大值是为了防止 exp 溢出。eager 模式下这个公式拆成 5 个 kernel:求 max、减 max、exp、求和、除。每个 kernel 都要读或写和 S 一样大的张量,5 个加起来一共读写 8 次,2048 MiB(图 3-4 里 max 到 ÷Σ 这 5 个 kernel 进出显存的箭头)。

FLOPs 和耗时因此对不上:softmax 的运算量只有 PV 的 1/5,耗时却是 PV 的 6 倍(图 3-6)。

矩阵乘逐元素FLOPsGPU 时间(ms)QKᵀQKᵀ:8.6e9 FLOPs(一层)8.6e9QKᵀ:0.3 ms(一层)0.30÷√d + mask÷√d + mask:6.7e7 FLOPs(一层)6.7e7÷√d + mask:0.75 ms(一层)0.75softmaxsoftmax:1.8e9 FLOPs(一层)1.8e9softmax:1.43 ms(一层)1.43PVPV:8.6e9 FLOPs(一层)8.6e9PV:0.23 ms(一层)0.23
图 3-6 一层 attention 里各 op 的 FLOPs 与实测 GPU 时间(medium,seq 1024)。

S、P 的读写量随 seq² 增长,Linear 只随 seq 线性增长。seq 从 256 增加到 1024,attention 占前向时间的比例从 10% 涨到 46%,多出来的几乎全是 softmax、除以 √d、mask 这类只搬数据的 op(图 3-7)。这就是第 2 节那 40% 里随 seq 涨得最快的部分。

softmaxscores(QKᵀ、÷√d、mask)PVattention 合计0%10%20%30%40%50%seq 256seq 512seq 1024占 forward GPU 时间(medium,24 层)seq 256:合计 2.3 ms,占 10%seq 512:合计 10.2 ms,占 22%seq 1024:合计 64.9 ms,占 46%合计 46%seq 256:softmax 1.0 ms,占 4%seq 512:softmax 5.0 ms,占 11%seq 1024:softmax 34.3 ms,占 24%softmax 24%seq 256:scores 0.9 ms,占 4%seq 512:scores 3.6 ms,占 8%seq 1024:scores 25.1 ms,占 18%scores 18%seq 256:PV 0.4 ms,占 2%seq 512:PV 1.6 ms,占 3%seq 1024:PV 5.5 ms,占 4%PV 4%
图 3-7 attention 三段占 forward GPU 时间的比例随 seq 变化(medium)。

要少搬,就得把几步合进一个 kernel,中间结果留在片上:融合的 softmax 只读一次 S、写一次 P;FlashAttention 更进一步,S、P 根本不写回显存。

显存上,表 3-3 是一层 attention 单独测的:新占显存的主要是两个 seq × seq 张量。放到整层 Transformer block 上也是这样,只是 torch.compile 之后存下的两个 seq × seq 张量换成了 S 和 P。xl 的一层(RMSNorm 这类中间量已经被省掉)一共要为反向存 3655 MiB,其中一半以上是 S 和 P(图 3-8)。这组测量用的是 16 头,S、P 各 1 GiB;标准 xl 是 32 头,S、P 还要再大一倍。

S、P([b, h, s, s]):2048 MiB,56.0%FFN 中间量([b, s, d_ff]):960 MiB,26.3%[b, s, d] 级张量(x、norm 输出、Q、K、V 等):640 MiB,17.5%其他(mask、RoPE、softmax 统计量):7 MiB,0.2%3655 MiBxl 一层S、P[b, h, s, s]2048 MiB · 56.0%FFN 中间量[b, s, d_ff]960 MiB · 26.3%[b, s, d] 级张量x、norm 输出、Q、K、V 等640 MiB · 17.5%其他mask、RoPE、softmax 统计量7 MiB · 0.2%
图 3-8 xl 一层为反向存的张量(batch 4,seq 2048,16 头,torch.compile 后用 saved_tensors_hooks 实测)。

xl 有 32 层,按 16 头算加起来也有 114 GiB,是 5090 显存的三倍多。而且只有 S、P 随 seq² 增长:32 头时,同一个 [b, h, s, s] 张量在 seq 128 时是 8 MiB,seq 2048 时是 2 GiB,是残差流上一个 [b, s, d] 张量的 25 倍。

图 3-9 是 xl 一步的显存时间线。seq 2048 只跑前向时,每层 attention 都让显存冲高约 8 GiB,算完再落回去,32 层都能跑完;加上反向后,每层的 S、P 都得留下,第 1 层就多占约 4.7 GiB,到第 2 层就 OOM 了。

seq 128 · 纯前向0102030权重 12.8 GiB峰值 12.90 GiB峰值 12.90 GiB10110 次分配 / 释放seq 2048 · 纯前向0102030权重 12.8 GiB峰值 21.38 GiB峰值 21.38 GiB10121 次分配 / 释放seq 128 · full step0102030权重 12.8 GiB峰值 28.57 GiB峰值 28.57 GiB反向optimizerOOM13354 次分配 / 释放seq 2048 · 前向 + 反向0102030权重 12.8 GiB峰值 25.96 GiB峰值 25.96 GiBOOM195 次分配 / 释放显存(GiB)
图 3-9 xl(batch 4,32 头)一步的显存时间线,横轴是分配 / 释放的次序。

4 优化:bf16 和 checkpoint 都差一口气
#

要减小 A 有两种办法:把每个张量存得小一点(bf16),或者少存一些、反向时重算(checkpoint)。

4.1 bf16:快了,省得不多
#

autocast 只把矩阵乘的输入换成 bf16,权重、梯度和 Adam 状态还是 fp32。速度提升很明显,前向快了 1.9–2.3 倍(图 4-1):矩阵乘换到 bf16 的 Tensor core 上,峰值从 1.05e14 翻倍到 2.1e14,要搬的字节也少了一半。显存只省了 18–21%:W、G 和 Adam 状态大小不变;A 也没有减半,因为 norm、softmax、残差和 loss 还在 fp32 下算(这些累加在 bf16 下不准,bf16 只有 7 位尾数,把 0.01 累加 1000 次只能得到 4.0),反向还要多存一份 bf16 的权重副本。存下的张量还是那些,只是一部分从 4 字节变成了 2 字节。

前向反向加速比(bf16 相对 fp32)0×1×2×small 前向:1.87×1.87small 反向:1.69×1.69smallmedium 前向:2.05×2.05medium 反向:1.79×1.79mediumlarge 前向:2.30×2.30large 反向:1.87×1.87largefp32bf16前向 + 反向的峰值显存(GiB)01020small fp32:4.08 GiBsmall bf16:3.18 GiB−21%smallmedium fp32:10.58 GiBmedium bf16:8.36 GiB−21%mediumlarge fp32:20.28 GiBlarge bf16:16.61 GiB−18%large
图 4-1 bf16 autocast 相对 fp32(前向 + 反向,batch 4,seq 512):左边是加速比,右边是峰值显存。

4.2 Checkpoint:省显存,多一遍前向
#

checkpoint 和 3.1 节里融合 RMSNorm 的做法一样,只是从一个 op 扩大到几层:前向只存每段的入口(entry,xl、seq 2048 时 80 MiB),反向走到这一段时,用入口把这段的前向重跑一遍,用完就释放。4 层 xl block 每 2 层设一个 checkpoint,峰值就从 4 × 3655 MiB = 14.6 GiB 降到「2 个 entry 加一段」的 7.5 GiB(图 4-2)。

全部存下峰值 14.6 GiB每 2 层一段峰值 7.5 GiBx0L13655 MiB含 x0L23655 MiB含 x1L33655 MiB含 x2L43655 MiB含 x3y4 份同时留到反向x0entry 80 MiBx2entry 80 MiBcheckpoint[L1 L2]2 × 3655 MiB反向时用 x0 重算,用完即丢checkpoint[L3 L4]2 × 3655 MiB反向时用 x2 重算,用完即丢y反向时一次只重算一段
图 4-2 4 层 xl block 有无 checkpoint:钢蓝框一直占到反向,浅蓝虚线框在反向时用 entry 重算、用完即丢。

代价是整个网络要多跑一遍前向。xl 在 seq 2048 下光参数加梯度就要 25.4 GiB,放不下 activation,所以我在 large 上扫了每段放几层(图 4-3)。不管怎么切,step 都是 302–313 ms,比不用 checkpoint 的 236 ms 慢 28–33%,正好对应一步从 3F 变成 4F(F 是一次前向,反向约 2F)。显存则是切得越细越省,每层一个 checkpoint 时最低,7.8 GiB,不用时是 15.0 GiB。

checkpoint不 checkpointstep 时间(ms)0100200300不 checkpoint 236 ms每段 1 层:302 ms每段 2 层:306 ms每段 3 层:311 ms每段 4 层:311 ms每段 6 层:312 ms每段 9 层:309 ms每段 12 层:312 ms每段 18 层:313 ms每段 36 层:313 ms123469121836每段几层(对数轴)峰值显存(GiB)0481216不 checkpoint 15.0 GiB每段 1 层:7.8 GiB每段 2 层:8.0 GiB每段 3 层:8.2 GiB每段 4 层:8.4 GiB每段 6 层:8.8 GiB每段 9 层:9.5 GiB每段 12 层:10.1 GiB每段 18 层:11.3 GiB每段 36 层:15.0 GiB123469121836每段几层(对数轴)
图 4-3 checkpoint 段长扫描(large,batch 1,seq 1024,前向 + 反向,fp32 eager)。

切得越细越省,是因为入口很小。设共 $L$ 层、每段 $e$ 层、入口大小 $a$、一层的 A 为 $A_1$,峰值约为 $\frac{L}{e}a + e\,A_1$。只要所有入口加起来比一层的 A 小(这里 36 × 5 MiB = 180 MiB < 220 MiB),$e = 1$ 就是最优。

但 checkpoint 解决不了 S、P。重算到某一层时,这一层的 S、P 照样要完整写进显存再读出来,seq 2048 时每层约 8 GiB 的峰值还在。


5 小结:问题都在 S、P
#

时间和显存的问题最后都落在 S、P 上。时间上,它们让 attention 成了 memory-bound,seq 1024 时占前向将近一半的时间;显存上,它们占一层 saved tensors 的一半以上,xl 在 seq 2048 时第 2 层就 OOM。bf16 只能把它们存小一点,checkpoint 只能推迟它们出现,都没能让它们离开显存。

要解决,得把 RMSNorm 的做法用到 attention 上:把 QKᵀ、softmax、PV 合进一个 kernel,分块在片上算完,S、P 不写回显存,反向需要时再重算。下一篇 FlashAttention 1–4 讲的就是这个。


References
#

  1. Williams, Waterman, Patterson. Roofline: An Insightful Visual Performance Model for Multicore Architectures. Communications of the ACM, 2009. 第 1 节的 roofline 模型。
  2. NVIDIA. NVIDIA RTX Blackwell GPU Architecture(白皮书)。RTX 5090 的峰值算力和显存带宽。
  3. Chowdhery et al. PaLM: Scaling Language Modeling with Pathways. arXiv:2204.02311, 2022. MFU 的定义。
  4. Databricks. LLM Inference Performance Engineering: Best Practices, 2023. MBU 的定义。
  5. Kaplan et al. Scaling Laws for Neural Language Models. arXiv:2001.08361, 2020. 每 token 训练 FLOPs ≈ 6N 的估算。
  6. Vaswani et al. Attention Is All You Need. NeurIPS 2017.
  7. Zhang, Sennrich. Root Mean Square Layer Normalization. NeurIPS 2019.
  8. Loshchilov, Hutter. Decoupled Weight Decay Regularization. ICLR 2019. AdamW。
  9. Micikevicius et al. Mixed Precision Training. ICLR 2018. 第 4.1 节的混合精度。
  10. Chen, Xu, Zhang, Guestrin. Training Deep Nets with Sublinear Memory Cost. arXiv:1604.06174, 2016. 第 4.2 节的 activation checkpoint。
  11. PyTorch 文档:Hooks for saved tensors(saved_tensors_hooks)、torch.profiler、torch.compile、Automatic Mixed Precision、torch.utils.checkpoint。
  12. Stanford CS336. Assignment 2: Systems, 2025. 计时和显存实验的设置。
  13. Dao, Fu, Ermon, Rudra, Ré. FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness. NeurIPS 2022. 下一篇的主题。
GPU training optimization - 这篇文章属于一个选集。
→ 谁偷走了 5090 的算力和显存:一步 Transformer 训练的 roofline 侦查 (本文)

Noviorlu喵

多伦多大学计算机工程。热爱计算机图形学、渲染和 AI/ML。

留言

用 GitHub 账号登录后可以留言。