训练加速实战:Flash Attention、Gradient Checkpointing 与数据流水线

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 的基础上做了两个额外优化:

  1. 减少非矩阵乘法操作:v1 里还有一些 rescaling 和 masking 操作是在单独的 CUDA kernel 里做的------每个 kernel launch 都有开销。v2 把这些操作融合进主循环,一个 kernel 搞定。
  2. 调整循环顺序:把外层循环从"按 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。

相关推荐
Loveyourself2 小时前
Claude Code Memory 总体系核心代码逐行解析
面试·agent
王林不想说话2 小时前
ES6 到 ES2026 全面进阶指南:新特性、原理、示例与工程落地
前端·javascript·面试
阿黎梨梨2 小时前
Next.js 全栈开发:从 SPA 的痛点到 SSR 的破局之道
前端·后端
名字还没想好☜2 小时前
Go 的 unsafe.Pointer 实战:零拷贝 []byte↔string 转换与三条铁律
开发语言·后端·golang·go·unsafe
老孙讲技术2 小时前
周末档期厨房爆单,带宽账单也爆了——不是直播难,是主码流 + always 计划在烧钱
后端·物联网
黄敬峰2 小时前
TypeScript 必考题:type 和 interface 的区别一次搞懂
面试
老孙讲技术2 小时前
校园开放日前夜才说要「全校透明」?我用轻应用把 9 路教室预览和回放嵌进了校园后台
后端·物联网
Python私教3 小时前
AI Agent 可观测性实战:从 correlationId 到失败时间线
人工智能·后端
用户7791666846543 小时前
给 Agent 开权限:身份不能进模型的上下文
后端