AI 加速器系列 · 第 3 篇
前面两篇解决了"显存不够"和"多卡通信"的问题。但这还不够快------真正的性能杀手藏在 Attention 计算和 I/O 里。
先看一组数字:一个 70B 参数的模型做一次标准 Attention 计算,如果序列长度是 4096,光是 QK^T 矩阵就有 4096 x 4096 = 16M 个元素。这还不算最糟的------最糟的是这个矩阵要频繁在 HBM 和计算单元之间来回搬运。一趟搬运的时间够 GPU 做几十次矩阵乘法。
这一篇拆三个训练加速技术:Flash Attention(让 Attention 不再被显存带宽卡住)、Gradient Checkpointing(用 30% 算力换 60% 显存)、DataLoader 优化(不让 GPU 闲着等数据)。
一、Flash Attention:把 Attention 的计算复杂度"骗"过去
真正的瓶颈不是 FLOPs,是显存带宽
很多人说 Attention 的瓶颈是 O(N^2) 的计算量------序列长度翻倍,计算量翻四倍。这句话对了一半。
O(N^2) 的 FLOPs 确实多,但 GPU 最擅长的就是算矩阵乘法------一个 A100 的 Tensor Core 每秒能做 312 TFLOPS 的 FP16 矩阵乘加。真正拖慢 Attention 的不是"算不动",而是"数据送不到计算单元手里"。
GPU 的存储分两层:
bash
┌─────────────────────────────────────────────────┐
│ HBM (High Bandwidth Memory) │
│ 80GB · 2 TB/s 带宽 │
│ 距离计算单元:远 │
│ 作用:存模型参数、中间激活值、优化器状态 │
├─────────────────────────────────────────────────┤
│ SRAM (Static RAM, on-chip) │
│ ~20MB 每 SM · 19 TB/s 带宽 │
│ 108 个 SM,分布在 GPU 芯片上 │
│ 距离计算单元:紧挨着 Tensor Core │
│ 作用:计算时的"草稿纸" │
└─────────────────────────────────────────────────┘
关键数字:HBM 有 80GB 但带宽"只有"2TB/s。SRAM 只有 20MB 但带宽高达 19TB/s------快了将近 10 倍。问题是 20MB 连一个 batch 的激活值都装不下。
标准 Attention 是怎么浪费带宽的
标准 Attention 的计算步骤:
css
1. S = Q × K^T → [N, N] 矩阵,写回 HBM
2. P = Softmax(S) → [N, N] 矩阵,从 HBM 读 S,算完写回 HBM
3. O = P × V → [N, d] 矩阵,从 HBM 读 P 和 V,算完写回 HBM
每一步都有一趟 HBM 读写。对于 N=4096、d=128(d_model / num_heads),这个 N, N 矩阵是 16M 个元素,FP16 下是 32MB------看起来不大。但问题在于:
- 这个矩阵不能一直留在 SRAM 里(太大了,20MB 装不下)
- 所以算完一步就得写回 HBM,下一步再从 HBM 读出来
- 一个 Transformer Block 有几十层,每层都要这么玩一次
- 整个训练过程,Attention 的 HBM 访问量达到 O(N^2) 级别
HBM 读写才是真正的瓶颈,不是计算量。
Flash Attention:在 SRAM 里原地算完
Flash Attention 的核心思路:别把中间矩阵写回 HBM。把 Q、K、V 切成小块,在 SRAM 里逐块计算,最终只把结果矩阵 O 写回 HBM。
具体做法分两步------Tiling 和 Recomputation。
Tiling(分块计算):
ini
标准做法:
整块 Q[N,d] × K^T[d,N] → 整块 S[N,N] → 写 HBM
Flash Attention:
把 Q 切成 Q1, Q2, ..., QT (每块 [B_r, d])
把 K、V 切成 K1, K2, ..., KT (每块 [B_c, d])
对每一块 Qi:
for each Kj, Vj: ← 这个循环在 SRAM 里跑
S_ij = Qi × Kj^T ← [B_r, B_c],小矩阵,留 SRAM
P_ij = Softmax(S_ij) ← 就地计算,不写 HBM
O_i += P_ij × Vj ← 累加到输出
结束循环
把 O_i 写回 HBM ← 只有这一次写 HBM
最终:只写 O[N,d] 回 HBM,中间所有 [N,N] 矩阵没离开过 SRAM
Recomputation(需要时重算,不要从 HBM 读):
Softmax 有个麻烦:它的分母是整行的 sum(exp(x))。分块计算时,每个块只知道自己那一部分------没法直接算出正确的 Softmax。标准做法是先把所有块的 exp 值写回 HBM,再读回来算分母------原地踏步。
Flash Attention 的做法:不写 HBM。在 SRAM 里维护一个 running max 和 running sum,每处理一个新块就更新这两个统计量,用代数变形保证最终的 Softmax 等价于全局计算。这个过程中如果某个块需要重新算 exp 值------就在 SRAM 里重算,比从 HBM 读回来还快。
效果:
scss
标准 Attention Flash Attention 改善
HBM 读/写量 O(N^2) O(N) 10-20x 减少
显存占用 O(N^2) O(N) 可训练更长序列
实际 wall-clock 基准 2-4x 加速 训练吞吐翻倍
Flash Attention 2:榨干最后一点性能
Flash Attention 2(2023)在 v1 的基础上做了两个额外优化:
- 减少非矩阵乘法操作:v1 里还有一些 rescaling 和 masking 操作是在单独的 CUDA kernel 里做的------每个 kernel launch 都有开销。v2 把这些操作融合进主循环,一个 kernel 搞定。
- 调整循环顺序:把外层循环从"按 Q 的分块"改为"按 K/V 的分块",让同一个 K/V 块被所有 Q 块复用------减少 SRAM 重载。
结果是 v1 的 2x 再加速,合计相比标准 Attention 可达到 4-8x。
PyTorch 里怎么用
PyTorch 2.0+ 内置了 Flash Attention 作为 scaled_dot_product_attention 的后端:
python
import torch
import torch.nn.functional as F
# 方法 1:直接用 PyTorch 的 SDPA(自动选择最优后端)
# PyTorch 会根据输入 shape、dtype、硬件自动选择 Flash Attention / Memory Efficient Attention / 标准实现
output = F.scaled_dot_product_attention(query, key, value, attn_mask=mask)
# 方法 2:torch.compile 自动融合
model = torch.compile(model)
# compile 会自动把 Attention 调用替换为高效的 F.scaled_dot_product_attention
# 方法 3:强制指定后端(调试用)
with torch.backends.cuda.sdp_kernel(
enable_flash=True,
enable_math=False,
enable_mem_efficient=False
):
output = F.scaled_dot_product_attention(q, k, v)
简单说:用 torch.compile() 包装你的模型,或者用 F.scaled_dot_product_attention 替换手写的 Attention------性能提升是自动的,不用改模型结构。
二、Gradient Checkpointing:拿 30% 算力换 60% 显存
反向传播为什么需要前向激活值
神经网络训练的核心是链式法则:
ini
loss = f4( f3( f2( f1(x) ) ) )
反向传播:
grad_f4 = ∂loss/∂f4 ← 需要 f3 的输出
grad_f3 = grad_f4 × ∂f4/∂f3 ← 需要 f3 的输入(= f2 的输出)和 grad_f4
grad_f2 = grad_f3 × ∂f3/∂f2 ← 需要 f2 的输出
...
计算每一层的梯度时,需要这一层前向计算的输入值(即上一层的输出------称为"激活值")。所以标准训练的流程是:
markdown
前向:x → Layer1 → a1 → Layer2 → a2 → ... → LayerL → aL → loss
存储 a1, a2, ..., aL-1 在显存里
反向:从 loss 开始,逐层读回激活值 a_i,算梯度
所有激活值留着不动------显存里最大的开销就是它们,不是模型参数。
问题:Transformer 的激活值有多大
对于一个 L 层的 Transformer,训练一个 batch 的显存占用:
scss
模型参数: 70B × 2 bytes (FP16) ≈ 140 GB ← 用 FSDP 切分后可以降低
优化器状态: 70B × 12 bytes (Adam) ≈ 840 GB ← 同样用 FSDP 切分
激活值: batch × seq_len × d_model × L × (几十个中间结果)
对于 batch=8, seq_len=4096, d_model=8192, L=80:
yaml
激活值 ≈ 8 × 4096 × 8192 × 80 × ~30 (中间变量数) × 2 bytes
≈ 8 × 4096 × 8192 × 80 × 30 × 2
≈ 1.2 TB ← 这才是一个 batch 的激活值
不切分或者不做任何优化,一个 batch 就能把 8 张 A100 全塞满。这就是为什么大模型训练经常 OOM------不是参数量大,是激活值太大了。
Checkpointing 的思路:不要全存,需要时重算
markdown
标准做法:
前向:Layer1 → Layer2 → Layer3 → ... → LayerL
存 a1 存 a2 存 a3 存 aL
反向:直接用 a_i 算梯度 ← 显存炸了
Checkpointing:
前向:Layer1 → Layer2 → Layer3 → ... → LayerL
存 a1 ← 只保留"检查点"
不存 不存 不存
反向:算 LayerL 的梯度 → 需要 a_{L-1} → 没有
→ 从最近的检查点 a1 重新前向算一遍
→ Layer1 → Layer2 → ... → Layer_{L-1} ← 重算
→ 拿到 a_{L-1},算梯度
→ 继续反向 Layer_{L-1},又需要 a_{L-2} → 又重算
一句话:只在某些层保留激活值("检查点"),反向传播时需要哪个层的激活值,就从最近的检查点重新前向计算一遍。空间换时间的反向操作------用时间(30% 额外前向计算)换空间(60% 显存节省)。
实践中用在哪几层
不需要对每一层都做 Checkpointing------Transformer 里不同层对显存的贡献差别很大:
erlang
一个 Transformer Block 的激活值分布(大致):
Attention:60-70%(QKV 投影 + Softmax 中间结果 + Value 聚合)
FFN: 20-25%(两个大矩阵乘法,中间隐藏维度 4× d_model)
LayerNorm:5-10%(就两个向量,几乎可以忽略)
策略:只对 Attention 和 FFN 做 Checkpointing,LayerNorm 不做。
python
import torch
import torch.utils.checkpoint as checkpoint
from torch import nn
class TransformerBlock(nn.Module):
def __init__(self, d_model, n_heads, d_ff):
super().__init__()
self.ln1 = nn.LayerNorm(d_model)
self.attention = nn.MultiheadAttention(d_model, n_heads, batch_first=True)
self.ln2 = nn.LayerNorm(d_model)
self.ffn = nn.Sequential(
nn.Linear(d_model, d_ff),
nn.GELU(),
nn.Linear(d_ff, d_model),
)
def forward(self, x):
# LayerNorm 不做 Checkpointing------开销可以忽略
h = self.ln1(x)
# Attention + 残差连接:包在 checkpoint 里
attn_out = checkpoint.checkpoint(
lambda h, x: x + self.attention(h, h, h, need_weights=False)[0],
h, x,
use_reentrant=False # PyTorch 1.11+ 推荐
)
# FFN + 残差连接:包在 checkpoint 里
h2 = self.ln2(attn_out)
output = checkpoint.checkpoint(
lambda h2, attn_out: attn_out + self.ffn(h2),
h2, attn_out,
use_reentrant=False
)
return output
要点:
use_reentrant=False(PyTorch 1.11+):不重入------前向的时候不会保存中间变量,反向时从头重跑。显存节省最大。layer norm不包 checkpoint:LayerNorm 的激活值极小(两个 d_model 向量),不值得花时间重算。- 残差连接也包进去:残差的分支(identity)也需要作为激活值------包在一起可以免存。
显存节省的量化估算
erlang
无 Checkpointing 有 Checkpointing 节省
一个 Transformer Block:
Attention ~10 MB ~2 MB (只存输出)
FFN ~4 MB ~1 MB
LayerNorm ~0.1 MB ~0.1 MB
──────────────────────────────────────────────────
per-block ~14 MB ~3 MB ~78%
总(80 layers) ~1.1 GB ~240 MB ~78%
实际场景中还要考虑 batch size 和 seq_len 的放大效应:
ini
batch=8, seq_len=4096:
无 Checkpointing:~9 GB 激活值
有 Checkpointing:~3.5 GB(节省 ~60%)
额外前向计算:~30% 训练时间增加
**30% 的训练时间换来 60% 的显存节省------这把算力换显存,在 GPU 显存是硬约束的场景下非常划算。**多出来的显存可以用来增大 batch size、增大模型宽度、或者训练更长序列------这些带来的模型质量提升远超过 30% 的额外训练时间。
三、DataLoader:别让 GPU 等 CPU
GPU 为什么闲着
真实训练中 GPU 的时间线是这样的:
ini
GPU:
|■■ batch 1 ■■| |■■ batch 2 ■■| |■■ batch 3 ■■|
CPU:
|██████ 准备 batch 1 ██████|██████ 准备 batch 2 ██████|██████ 准备 batch 3 ██████|
时间 →
0s 0.1s 0.5s
GPU 算一个 batch:0.1s
CPU 准备一个 batch:0.5s(读磁盘 + 解码 JPEG + resize + normalize + 拼 batch)
GPU 利用率:0.1 / 0.5 = 20% ← 80% 的时间在等 CPU
不是 GPU 慢,是 CPU 来不及喂。CPU 要做的数据预处理比很多人想象的更重:
scss
一个训练 step 的数据预处理管线:
Raw File → 读磁盘 (I/O) → 解码 (JPEG/PNG/Audio) → Resize/Crop →
Normalize → ToTensor → 拼成 batch → CPU → GPU (PCIe 传输)
每一步都是瓶颈------但最慢的通常是磁盘 I/O 和图像解码。
PyTorch DataLoader 的核心参数
python
from torch.utils.data import DataLoader
loader = DataLoader(
dataset,
batch_size=32,
num_workers=8, # ← 开 8 个子进程并行准备数据
prefetch_factor=4, # ← 每个 worker 提前准备 4 个 batch
pin_memory=True, # ← 把数据锁在 page-locked memory,加速 CPU→GPU 传输
persistent_workers=True, # ← worker 进程不销毁(避免每个 epoch 重新 fork)
drop_last=True, # ← 丢弃最后一个不完整的 batch
)
逐个说明:
num_workers ------ 并行加载的进程数。经验值:num_workers = 4 × num_gpus 或 CPU 核心数的一半。太少 GPU 等数据;太多进程切换开销吃掉收益。判定方式:看 nvidia-smi 里 GPU 利用率有没有波动------波动大说明 num_workers 不够。
markdown
num_workers=0(主进程加载):
GPU: |■■ batch ■■|______________|■■ batch ■■|______________|
CPU: 加载中 加载中
num_workers=8:
GPU: |■■ batch ■■|■■ batch ■■|■■ batch ■■|■■ batch ■■|
CPU: 加载中 (并行 8 个) ← 下一批已经在准备好了
prefetch_factor ------ 每个 worker 提前预取多少个 batch 放在内存里。默认是 2。如果你有一个很快的 GPU 和一个很慢的磁盘,把这个值拉到 4-8------让 CPU 比 GPU 快一步,保证 GPU 永远不空等。
pin_memory=True ------ 把加载好的数据分配在"锁页内存"(page-locked memory)里。普通内存页可以被 OS swap 出去;锁页内存不可以。好处是 DMA(Direct Memory Access)可以直接从这块内存往 GPU 拷数据------不需要 CPU 在中间做一次拷贝。效果是 CPU→GPU 传输快 2-3 倍。
persistent_workers=True ------ 默认情况下每个 epoch 结束后 worker 进程会被销毁再重新创建。对于大模型训练(一个 epoch 几个小时),这个开销无关紧要。但对于小数据集多 epoch 的场景,fork 进程的开销会累积。开了这个参数,worker 一直活着。
调参口诀
ini
GPU 利用率 < 90% 且波动大 → num_workers 不够,往上加
训练刚开始 GPU 利用率低 → prefetch_factor 太小,加大到 4-8
CPU 使用率很高但 GPU 利用率不高 → 数据预处理太重,上 DALI
每个 epoch 开头特别慢 → 开 persistent_workers=True
NVIDIA DALI:GPU 上直接做预处理
PyTorch DataLoader 的预处理跑在 CPU 上------JPEG 解码、Resize、Crop 这些操作。CPU 的并行度有限,而且数据还要跨 PCIe 总线从 CPU 搬到 GPU。
DALI(Data Loading Library)的思路:把整个数据预处理管线搬到 GPU 上跑。
scss
标准流程:
磁盘 → CPU(解码+resize+normalize) → PCIe 传输 → GPU(训练)
DALI 流程:
磁盘 → GPU(解码+resize+normalize+训练) ← 全在 GPU 上
GPU 做图像解码远快于 CPU------一个 A100 能同时并行解码几百张 JPEG。而且数据解码完后直接就在 GPU 显存里了------不需要跨 PCIe 传输。
python
from nvidia.dali.pipeline import pipeline_def
import nvidia.dali.fn as fn
from nvidia.dali.plugin.pytorch import DALIGenericIterator
@pipeline_def(batch_size=256, num_threads=4, device_id=0)
def train_pipeline():
images, labels = fn.readers.file(
file_root="/data/imagenet/train",
random_shuffle=True,
name="Reader"
)
images = fn.decoders.image(images, device="mixed") # GPU 解码
images = fn.resize(images, resize_x=224, resize_y=224)
images = fn.crop_mirror_normalize(
images,
mean=[0.485 * 255, 0.456 * 255, 0.406 * 255],
std=[0.229 * 255, 0.224 * 255, 0.225 * 255],
)
return images, labels
pipeline = train_pipeline()
pipeline.build()
loader = DALIGenericIterator(pipeline, ["data", "label"])
DALI 适用于:大批量图像训练(ImageNet、COCO)、视频解码训练、音频预处理。对于 NLP 任务(文本已经够轻了),DALI 的收益不大------PyTorch DataLoader 调好参数就够用。
数据格式的选择
ini
原始 JPEG 文件:
小文件巨多(百万级),随机读取慢,文件系统 inode 压力大
→ 训练时 60% 时间在等磁盘寻道
TFRecord(TensorFlow):
多个 JPEG 打包成二进制序列文件,顺序读快
→ 但解码还是要在 CPU 上做
Parquet / Arrow:
列式存储,零拷贝读取
→ 对于结构化数据(文本 token、metadata)极快
→ 配合 Arrow 内存格式:数据在磁盘上的 layout = 数据在内存里的 layout
→ 不需要反序列化,mmap 直接当内存用
WebDataset(tar 分片):
把图片打包成 .tar 文件,每个文件几百 MB
→ 顺序读 I/O 很快,不伤文件系统
→ 社区推荐的标准做法
典型的大规模训练数据管线:
markdown
原始图片/文本
│
▼
预处理 + Tokenize(一次性)
│
▼
WebDataset / Parquet 分片(列存 + 零拷贝)
│
▼
PyTorch DataLoader / DALI(并行加载 + GPU 预处理)
│
▼
训练循环(GPU 不空转)
一句话总结
Flash Attention 解决 Attention 被 HBM 带宽卡住 → 把计算留在 SRAM 里不用来回搬;Gradient Checkpointing 解决激活值占显存太大 → 扔掉的中间结果反向时原地重算;DataLoader 优化解决 GPU 等数据 → 预处理管线化、并行化、必要时搬到 GPU 上做。三者叠加:同样硬件、同样模型,训练吞吐翻 2-3 倍。
下一篇:混合精度训练与 FP8------怎么用一半的位宽,训练出同样的精度。BF16 为什么比 FP16 更适合训练,FP8 的 E4M3 和 E5M2 有什么区别,以及怎么在 PyTorch 里两行代码开启 AMP。