CS336 · 从零构建语言模型
Lecture 5 · GPU
源文件:lecture-05.md
Lecture 5 · GPU
CS336: Language Modeling from Scratch · Stanford · Spring 2025
📅 Apr 15 · 💻 不可执行讲义(PDF) · 讲者 Tatsunori Hashimoto
承上启下
上一讲(Lecture 4 · 混合专家模型 MoE):
- MoE 动机:同样 FLOPs 下塞更多参数——路由机制让每个 token 只激活少数专家,推理成本接近稠密小模型、容量接近稠密大模型
- 路由函数:top-k 门控、哈希路由、RL 学习的路由器;专家粒度(粗粒度 FFN vs 细粒度参数分片)与共享专家(shared experts)
- 训练挑战:负载均衡损失(load balancing loss)防止专家坍缩、router z-loss 稳定 logits、梯度裁剪与初始化技巧
- 并行与工程:expert parallelism、all-to-all 通信、upcycling(从稠密模型初始化)
- 真实案例:DeepSeek V1→V3(细粒度 MoE + 多 token 预测)、Mixtral 8x7B、Qwen MoE、OLMoE
本讲(GPU):
- MoE 把模型做大了,但训练/推理的瓶颈在硬件——为什么 A100/H100 这么贵?为什么同样的算法在 GPU 上快 100 倍?在 Lecture 2 · PyTorch 与资源核算 里我们学会了数 FLOPs 和内存,本讲解释这些数字如何落到真实硬件上
- 语言模型训练是计算密集 + 内存密集的双重游戏:要榨干 GPU 的算力(TFLOPS)、也要驯服内存墙(memory wall)
- 本讲目标:打开 GPU 黑箱,理解硬件如何决定代码快慢;学会用 roofline 模型预测性能上限;掌握让代码「跑满 GPU」的核心技巧
说明
GPU 是一台并行计算机器 + 分层内存系统;语言模型训练的核心矛盾是算力(TFLOPS)vs 带宽(GB/s);优化的目标是提高算术强度(arithmetic intensity)——每从内存搬一个字节就做更多计算:让数据在片上反复使用(tiling)、合并访问(coalescing)、融合算子(fusion)、必要时用重计算换内存(recomputation)。FlashAttention 就是这一切技巧的集大成。
1. GPU 硬件解剖:一台并行计算机器
1.1 为什么是 GPU?两个数字的对比
CPU 的设计哲学是快速串行执行:芯片面积大量花在复杂控制流、大缓存、分支预测、乱序执行上,目标是让单个线程尽可能快。GPU 的设计哲学是海量并行吞吐:把控制逻辑做到最简、用省下来的晶体管堆成千上万个简单核心,让它们同时做同一件事(SIMT),目标是让总吞吐量尽可能大。
| 硬件 | 算力(FP16/BF16) | 内存带宽 | 典型用途 |
|---|---|---|---|
| Intel Xeon(24 核 CPU) | ~1 TFLOPS | ~100 GB/s | 通用计算、复杂控制流 |
| NVIDIA H100 GPU | 989 TFLOPS(Tensor Core,稠密) | 3.35 TB/s(HBM3) | 大规模并行、矩阵运算 |
注意
989 TFLOPS 是 H100 SXM 的 BF16/FP16 Tensor Core 稠密算力。如果你的代码走的是普通 FP32 CUDA core 路径,峰值只有约 67 TFLOPS——差 15 倍。「用不用得上 Tensor Core」本身就是最大的优化项之一(需要低精度 dtype + 合适的矩阵形状)。
提示
算力差距约 1000 倍(989 vs ~1 TFLOPS),这就是训练语言模型必须用 GPU 的原因。但注意:带宽只快了约 30 倍(3.35 TB/s vs ~100 GB/s)。算力增长远快于带宽增长,意味着「喂数据」的速度越来越跟不上「算数据」的速度——这个失衡(内存墙)正是本讲一切优化技巧的根源。
graph TD
subgraph CPU["CPU 设计:少而强"]
C1["核心 1<br/>复杂控制单元"]:::cpu
C2["核心 2"]:::cpu
C3["..."]:::cpu
C4["核心 24"]:::cpu
CACHE["大 L3 缓存<br/>~100 MB"]:::cache
C1 --> CACHE
end
subgraph GPU["GPU 设计:多而简"]
SM1["SM 1<br/>128 个 FP32 核"]:::gpu
SM2["SM 2"]:::gpu
SM3["..."]:::gpu
SM4["SM 132<br/>全卡共 16,896 核"]:::gpu
HBM["HBM3 内存<br/>80 GB"]:::mem
SM1 --> HBM
end
CPU -->|"~1 TFLOPS<br/>少量线程"| GPU
GPU -->|"~1000 TFLOPS<br/>数万线程"| HBM
classDef cpu fill:#e1f5fe,stroke:#0277bd
classDef gpu fill:#c8e6c9,stroke:#2e7d32
classDef cache fill:#f3e5f5,stroke:#7b1fa2
classDef mem fill:#ffe0b2,stroke:#e65100
1.2 GPU 内部结构:SM 与 SIMT
NVIDIA GPU 架构层次(以 H100 为例):
- GPU 芯片(chip):整个物理设备
- 流多处理器(Streaming Multiprocessor, SM):GPU 的「核心单元」,H100 有 132 个 SM(本课程实验用的 A100 是 108 个)。一个 SM 是独立的调度与执行单元,有自己的寄存器堆、共享内存和 warp 调度器
- CUDA 核心(CUDA core):每个 SM 内有 128 个 FP32 核心,总共 132 × 128 = 16,896 个
- Tensor Core:专用矩阵乘加单元(每周期完成一个小矩阵块的乘加),每个 SM 配备 4 个(H100 是第四代)。矩阵乘法的峰值算力几乎全部来自 Tensor Core,CUDA core 只负责其余「杂活」
SIMT(Single Instruction, Multiple Threads)执行模型:程序员写「单个线程做什么」,硬件负责把同一段代码复制到海量线程上并行执行。
# CPU 风格:一个线程顺序执行
for i in range(1000000):
c[i] = a[i] + b[i] # 一次处理一个元素
# GPU 风格:启动 1000000 个线程,每个线程处理一个元素
def gpu_kernel(i): # 每个线程执行同样的代码
c[i] = a[i] + b[i] # 但 i 不同(线程索引)
# launch_kernel(gpu_kernel, num_threads=1000000)
Warp(线程束):GPU 的基本调度单位,32 个线程为一组,锁步(lockstep)执行同一条指令。硬件不会为每个线程单独取指、译码——一条指令喂 32 个线程,这正是 GPU 能省下控制逻辑面积的原因。
注意
如果同一个 warp 内的线程走了不同分支,GPU 无法「一条指令喂 32 个线程」,只能串行执行所有分支路径、把不满足条件的线程 mask 掉——两个分支就意味着并行度减半。
if thread_id % 2 == 0:
# 一半线程执行这里
else:
# 另一半执行这里 ← warp 内分支 = 性能减半
实践中应尽量让分支边界对齐 warp 边界(例如按 32 的倍数划分任务),或用无分支写法(torch.where、mask 乘法)代替 if-else。
1.3 内存层级:金字塔与延迟
GPU 的内存是分层的金字塔——越快的越小、越慢的越大。数字记不住不要紧,记住相邻层之间差一个数量级这个感觉:
| 层级 | 容量(H100) | 带宽 | 延迟 | 访问范围 |
|---|---|---|---|---|
| 寄存器(Register) | ~33 MB 全卡(256 KB / SM) | 极高(几乎不构成瓶颈) | ~1 周期 | 每个线程私有 |
| 共享内存(Shared Memory / SRAM) | 228 KB / SM | ~20-30 TB/s(全卡合计) | ~20-30 周期 | 同一 thread block 内共享 |
| L2 缓存(Cache) | 50 MB | ~10 TB/s 量级 | ~200 周期 | 全 GPU 共享、硬件自动管理 |
| 全局内存(Global / HBM3) | 80 GB | 3.35 TB/s | ~400-600 周期 | 所有线程可见 |
注意共享内存与 L2 的一个本质区别:共享内存由程序员显式管理(你决定放什么、什么时候放),L2 是硬件自动管理的缓存(你只能祈祷命中)。高性能 kernel 的核心手艺就是显式调度共享内存。
graph TB
REG["寄存器<br/>256 KB / SM<br/>~1 周期<br/>线程私有"]:::fast
SHMEM["共享内存 SRAM<br/>228 KB / SM<br/>~20 周期<br/>block 内共享"]:::fast
L2["L2 缓存<br/>50 MB<br/>~200 周期<br/>全 GPU 自动管理"]:::med
HBM["HBM3 全局内存<br/>80 GB<br/>3.35 TB/s<br/>~400+ 周期"]:::slow
REG --> SHMEM
SHMEM --> L2
L2 --> HBM
HBM -.慢 400 周期.-> L2
L2 -.中 200 周期.-> SHMEM
SHMEM -.快 20 周期.-> REG
classDef fast fill:#c8e6c9,stroke:#2e7d32
classDef med fill:#fff9c4,stroke:#f57f17
classDef slow fill:#ffcdd2,stroke:#c62828
提示
把 HBM 想成郊区的大仓库、SRAM 想成车间旁的料架、寄存器想成手边的工具台。优化的核心是让数据尽量待在快速的小内存里反复使用,而不是每算一步都跑一趟仓库——这叫 tiling(分块) 或 blocking。FlashAttention 的全部魔法都建立在这一点上。
2. 算力 vs 带宽:两种瓶颈、两种优化方向
2.1 两个极端:compute-bound vs memory-bound
衡量一个算子「性格」的关键指标是算术强度(Arithmetic Intensity, AI):
$$\text{AI} = \frac{\text{计算量(FLOPs)}}{\text{内存访问量(Bytes)}}$$
它回答的问题是:每从 HBM 搬运一个字节,能做多少次浮点运算? AI 高的算子「值回票价」(搬一次数据算很多次),AI 低的算子让算术单元一直在等数据。
Compute-bound(算力瓶颈)的例子:大矩阵乘法 $C = AB$,$A, B \in \mathbb{R}^{N \times N}$,$N = 8192$,BF16:
- 计算量:$2N^3 = 2 \times 8192^3 \approx 1.1 \times 10^{12}$ FLOPs(每个输出元素做 $N$ 次乘加,乘加算 2 FLOPs)
- 内存访问:读 $A$、$B$ 写 $C$,共 $3N^2 \times 2\ \text{B} \approx 402\ \text{MB}$
- 算术强度:$\text{AI} \approx 1.1 \times 10^{12} / 4.02 \times 10^{8} \approx 2700$ FLOPs/byte
Memory-bound(带宽瓶颈)的例子:逐元素加法 $C = A + B$,同样形状:
- 计算量:$N^2 \approx 6.7 \times 10^{7}$ FLOPs(每个元素加一次)
- 内存访问:仍然是 $3N^2 \times 2\ \text{B} \approx 402\ \text{MB}$(读 A、读 B、写 C)
- 算术强度:$\text{AI} \approx 0.17$ FLOPs/byte
同样的数据搬运量,矩阵乘法做了 16000 倍的计算——这就是为什么 GEMM 能跑满算力,而逐元素操作永远在等内存。
| 对比 | Compute-bound | Memory-bound |
|---|---|---|
| 特征 | 算术强度高(>295 FLOPs/byte) | 算术强度低(<10 FLOPs/byte) |
| 瓶颈 | TFLOPS 不够 | 带宽不够 |
| 优化方向 | Tensor Core、更低精度 | 减少内存访问、合并访问、算子融合 |
| 典型算子 | 矩阵乘法(GEMM)、卷积 | LayerNorm、Softmax、GELU、逐元素操作 |
注意
训练时间不是只由矩阵乘法决定——大量时间花在 memory-bound 的「小算子」上(激活函数、归一化、attention 的非 GEMM 部分)。GEMM 已经被 cuBLAS/Tensor Core 优化到接近极限,优化小算子(融合、减少 HBM 往返)才是普通人性能提升的主战场——这也是下一讲 Lecture 6 · Kernel 与 Triton 的主题。
2.2 Roofline 模型:性能上限的可视化
Roofline 模型(Williams et al., 2009) 用一张图同时刻画算力和带宽两个上限。给定算子的算术强度 $\text{AI}$,它能达到的性能上限是:
$$P = \min\left(\pi,\ \beta \cdot \text{AI}\right)$$
其中 $P$ 是可达到的性能(FLOPs/s),$\pi$ 是硬件峰值算力(H100 为 $989 \times 10^{12}$ FLOPs/s),$\beta$ 是内存带宽($3.35 \times 10^{12}$ B/s)。直觉:性能要么被「算的速度」封顶,要么被「喂数据的速度」封顶,取其小者。
性能(TFLOPS)
│
989 ├────────────────────────── 算力上限(H100 峰值 989 TFLOPS)
│ ╱
│ ╱ ← Roofline(屋顶线)
│ ╱
│ ╱
│ ╱ ← 斜线部分:带宽瓶颈区域
│ ╱ 性能 = 带宽 × AI
├───────────────────────────── AI(FLOPs/byte)
0 295 → AI* = 989 TFLOPS / 3.35 TB/s ≈ 295
↑
「屋脊点」:AI 再高也不会更快
屋脊点(ridge point):两条线的交点
$$\text{AI}^{*} = \frac{\pi}{\beta} = \frac{989 \times 10^{12}}{3.35 \times 10^{12}} \approx 295\ \text{FLOPs/byte}$$
- $\text{AI} < 295$:memory-bound,此时 $P = \beta \cdot \text{AI}$——提高 AI(融合、tiling)或带宽利用率才能加速
- $\text{AI} > 295$:compute-bound,此时 $P = \pi$——只能换更强的算术单元或降低精度
例子
BF16 加法的 $\text{AI} \approx 1/6$ FLOPs/byte(每 6 字节搬运做 1 次加法),代入 roofline:$P = 3.35 \times 10^{12} \times \frac{1}{6} \approx 0.56$ TFLOPS——不到峰值算力的 0.06%。无论怎么优化 kernel 本身,这个算子都不可能更快;唯一出路是把它融合进别的算子、彻底消灭这次内存往返。
提示
- 算子融合(fusion):把多个 memory-bound 算子合并、消灭中间结果的内存往返 → 点向右移(AI 提高)
- Tiling(分块):把数据切块放进共享内存反复使用 → 点向右移
- 低精度(BF16/FP8):每个数字节数减半 → AI 翻倍(点右移),同时 Tensor Core 峰值翻倍(屋顶抬高)
- 重计算(recomputation):不改变单个算子的 AI,而是用富余算力换稀缺内存,让更大的 batch/模型放得下
graph LR
A["算子 AI 低"]:::red --> F["融合 Fusion"]:::blue
F --> B["AI 提升<br/>减少内存往返"]:::green
C["数据反复从 HBM 加载"]:::red --> T["Tiling<br/>分块到 SRAM"]:::blue
T --> D["片上反复使用<br/>AI 提升"]:::green
E["FP32 4 字节"]:::red --> L["降精度 BF16"]:::blue
L --> G["2 字节 带宽减半<br/>Tensor Core 加速"]:::green
classDef red fill:#ffcdd2,stroke:#c62828
classDef blue fill:#e1f5fe,stroke:#0277bd
classDef green fill:#c8e6c9,stroke:#2e7d32
3. 低精度训练:FP16、BF16、FP8
3.1 浮点数格式回顾
IEEE 754 浮点数的位布局是「符号位 | 指数位 | 尾数位」:指数位决定动态范围(能表示多大/多小的数),尾数位决定精度(相邻两个可表示数之间的间隔)。
| 格式 | 符号 | 指数 | 尾数 | 总位数 | 动态范围 | 精度 |
|---|---|---|---|---|---|---|
| FP32 | 1 | 8 | 23 | 32 | ±3.4×10³⁸ | ~7 位十进制 |
| FP16 | 1 | 5 | 10 | 16 | ±6.5×10⁴ | ~3 位十进制 |
| BF16 | 1 | 8 | 7 | 16 | ±3.4×10³⁸(同 FP32) | ~2 位十进制 |
| FP8 (E4M3) | 1 | 4 | 3 | 8 | ±448 | ~1 位十进制 |
| FP8 (E5M2) | 1 | 5 | 2 | 8 | ±57,344 | ~0.5 位十进制 |
提示
BF16(Brain Float 16) 是 Google TPU 团队发明的格式——牺牲精度(尾数只有 7 位)换取与 FP32 完全相同的动态范围(指数 8 位)。深度学习里张量的数值跨越很多数量级(梯度可以是 1e-8,激活可以是 1e+3),溢出/下溢是灾难性的(loss 变 NaN 或梯度归零),而精度损失只是噪声——SGD 本来就在噪声中工作。所以「保范围、弃精度」是正确的取舍。额外好处:FP32 → BF16 只需截断低 16 位,硬件转换极其简单。
为什么低精度可行?
1. 梯度下降对精度不敏感:mini-batch 梯度本身就是带噪声的估计,尾数少几位引入的舍入误差远小于采样噪声
2. 动态范围比精度更重要:溢出会直接破坏训练,舍入不会
3. 硬件大幅加速:H100 Tensor Core 的 BF16 算力(989 TFLOPS)是 TF32 的 2 倍、是 FP8(1979 TFLOPS)的一半——精度每减半,峰值算力翻倍,同时内存占用和带宽消耗也减半
FP8 的两个变体分工:E4M3(范围小、精度高)用于前向的权重和激活;E5M2(范围大、精度低)用于反向的梯度——因为梯度的数值范围波动更剧烈。DeepSeek-V3 是首个大规模用 FP8 完成预训练的开源模型,需要精细的分块缩放(per-tile scaling)校准。
3.2 混合精度训练(Mixed Precision Training)
低精度不是「一刀切全换」,而是在不同环节用不同精度(Micikevicius et al., 2018):
- 前向 + 反向的矩阵乘法:用 BF16/FP16 计算(吃满 Tensor Core、省带宽)
- 权重主副本(master weights):用 FP32 存储——每步更新量 $\eta \cdot g$ 可能远小于权重本身,低精度下会被舍入吞掉($w + \delta = w$),FP32 主副本保证微小更新能被累积
- 梯度累加与优化器更新:在 FP32 下进行,避免大量小数相加的舍入误差累积
# 伪代码(PyTorch 自动混合精度)
model_fp32 = ... # FP32 主权重
optimizer = torch.optim.AdamW(model_fp32.parameters())
with torch.autocast(device_type='cuda', dtype=torch.bfloat16):
logits = model_fp32(input) # 自动转 BF16 计算
loss = criterion(logits, target)
loss.backward() # 梯度在 FP32 累积
optimizer.step() # FP32 更新
Loss scaling(损失缩放):FP16 的动态范围下限太高,很多小梯度会下溢变成 0。技巧是把损失放大 $s$ 倍再反向传播:
$$\tilde{\mathcal{L}} = s \cdot \mathcal{L} \quad \Rightarrow \quad \tilde{g} = s \cdot g$$
其中 $s$ 是缩放因子(如 512 或动态调整),$g$ 是真实梯度。放大后的梯度 $\tilde{g}$ 落回 FP16 可表示范围内,更新前再除以 $s$ 还原。BF16 动态范围与 FP32 相同,不需要 loss scaling——这是它成为默认选择的重要原因。
说明
现代训练默认用 BF16 混合精度(Llama 3、DeepSeek 早期版本都是)——比 FP16 稳定、免 loss scaling、硬件支持好。FP8 训练已被 DeepSeek-V3 验证可行,但需要精细的缩放管理,仍属前沿工程。推理侧则更激进:INT8/FP8/INT4 量化已是常态(见 Lecture 10 · 推理)。
4. 内存访问优化:Coalescing、Tiling、Fusion
4.1 内存合并(Memory Coalescing)
GPU 访问 HBM 是按事务(transaction)进行的——每次事务搬运一个连续的对齐内存块(缓存行 128 字节,即 32 个 FP32 或 64 个 FP16)。硬件会把一个 warp 内 32 个线程的访问请求合并成尽量少的事务。
合并访问(coalesced access):warp 内 32 个线程访问连续的 32 个元素(如 A[0..31])→ 合并成 1 次事务,带宽利用率 100%。
非合并访问(uncoalesced / strided access):warp 内线程访问间隔很大的地址(如 A[0], A[1000], A[2000], ...)→ 每个线程触发独立事务,最坏 32 次事务,每次搬 128 字节只用 4 字节——带宽利用率 3%。
# 好:合并访问(C 连续存储)
for i in range(N): # 每个线程处理一个 i
C[i] = A[i] + B[i] # warp 内 32 线程访问 C[0..31]、C[32..63]...
# 坏:非合并(按列访问行优先矩阵)
for i in range(M):
C[i] = A[i, 0] # A 是 (M, N) 行优先存储
# 访问 A[i,0] 的 stride = N → 每个线程隔 N 个元素
注意
- NumPy/PyTorch 默认行优先(row-major / C-order):
A[i,j]的地址 =base + i*N + j,同一行连续、同一列相隔 N A.T/A.transpose()只改变 stride 元信息,不搬动数据——转置后按「行」访问实际上是在按原来的列跳着读,可能从合并变成非合并- 症状:某个操作在
.contiguous()之后突然变快,通常就是这个原因
优化技巧:
1. 调整循环/线程映射:让「相邻线程」对应「相邻内存地址」(最内层维度)
2. Padding:矩阵宽度补齐到 32 的倍数,避免跨缓存行的错位访问
3. 共享内存中转:先按合并方式把数据搬进 shared memory,再在片上以任意模式读取(片上访问不受合并约束)
提示
HBM 像超市货架:拿东西必须整箱搬(128 字节事务)。32 个人(warp)如果买的是同一箱里的 32 件商品,搬一箱就够;如果每人要的商品散布在 32 个货架上,就得搬 32 箱、每箱只用一件。合并访问就是「让结伴的线程买同一箱货」。
4.2 Tiling(分块):让数据待在片上
核心思想:把大矩阵切成小块(tile),每块加载进共享内存后反复使用,把「每个数据从 HBM 读很多次」变成「读一次、片上用很多次」。
以矩阵乘法 $C = AB$($M \times K$ 乘 $K \times N$)为例。
朴素实现(每线程独立算一个输出):
# 每个线程计算 C[i, j]
C[i, j] = sum(A[i, k] * B[k, j] for k in range(K))
每个线程做 $2K$ FLOPs、从 HBM 读 $2K$ 个元素($K$ 个来自 $A$、$K$ 个来自 $B$,FP16 共 $4K$ 字节),算术强度只有
$$\text{AI}_{\text{naive}} = \frac{2K}{4K} = 0.5\ \text{FLOPs/byte}$$
对照屋脊点 295,这是灾难性的低——朴素 matmul 是彻底 memory-bound 的。
Tiling 实现:一个 thread block 负责 $C$ 的一个 $T \times T$ 块,沿 $K$ 维分步推进:
# 伪代码(简化,忽略边界情况),T = TILE_SIZE = 32
for t in range(0, K, T):
# 1. 协作加载:block 内所有线程合作把两个 tile 搬进共享内存(合并访问)
A_shared = A[i : i+T, t : t+T] # (32, 32)
B_shared = B[t : t+T, j : j+T] # (32, 32)
__syncthreads() # 同步,确保所有线程都加载完毕
# 2. 计算:每个线程用共享内存里的数据累加
for k in range(T):
C_local += A_shared[local_i, k] * B_shared[k, local_j]
__syncthreads()
# 写回 C[i, j] = C_local
效果分析:每个 tile 从 HBM 加载一次、被 block 内 $T$ 行/列的计算共用 $T$ 次,HBM 总流量降为原来的 $1/T$,算术强度提升 $T$ 倍:
$$\text{AI}_{\text{tiled}} \approx 0.5 \times T = 16\ \text{FLOPs/byte} \quad (T = 32)$$
实际的高性能 GEMM(cuBLAS/CUTLASS)还会叠加寄存器分块(register blocking)、更大的 tile、双缓冲流水线,把 AI 推到数百、逼近屋脊点——这正是矩阵乘法能跑满 Tensor Core 的原因。
提示
Tiling 的本质是复用:矩阵乘法天然「每个输入元素要参与 N 次计算」,朴素实现把这 N 次拆散到 N 次 HBM 读取;tiling 把它们聚拢到一次读取 + N 次片上访问。数据复用的机会是算法固有的,tiling 只是把这个机会从「L2 撞运气」变成「共享内存里明确兑现」。
graph TB
subgraph NO["无 Tiling:每个输出元素独立读"]
A1["读 A 的第 i 行<br/>共 K 次"] --> C1["C 的一个元素"]
B1["读 B 的第 j 列<br/>共 K 次"] --> C1
end
subgraph YES["有 Tiling:协作加载 + 片上复用"]
A2["A tile 加载 1 次<br/>到 SRAM"] --> SRAM["共享内存<br/>32×32 tile"]
B2["B tile 加载 1 次<br/>到 SRAM"] --> SRAM
SRAM --> C2["计算 32×32 个输出<br/>每个数据片上复用 32 次"]
end
style C1 fill:#ffcdd2,stroke:#c62828
style C2 fill:#c8e6c9,stroke:#2e7d32
style SRAM fill:#e1f5fe,stroke:#0277bd
4.3 算子融合(Operator Fusion)
问题:PyTorch 默认逐算子执行(eager mode)——每个算子是独立的 kernel,从 HBM 读输入、写输出,互相之间通过 HBM 传递中间结果。
# 三个算子,三次 HBM 往返
x = layer_norm(x) # 读 x,写 x_norm
x = gelu(x) # 读 x_norm,写 x_act
x = dropout(x) # 读 x_act,写 x_drop
# 总内存访问:6 × size(x)(三次读 + 三次写)
算子融合:把多个算子合并成一个 kernel,中间结果只存在寄存器/共享内存里、永不落 HBM。
# 融合后:一次读 x,一次写 x_out
x_out = fused_ln_gelu_dropout(x)
# 内存访问:2 × size(x) → 内存流量降为 1/3
对 memory-bound 算子,内存流量降为 $1/k$ 就意味着接近 $k$ 倍加速——计算量根本不是瓶颈。
适用场景:
- 逐元素操作链:GELU、ReLU、Dropout、残差加法等任意组合
- 小归约操作:Softmax、LayerNorm/RMSNorm(行内归约 + 逐元素变换)
- 注意力机制:QK 乘积 + softmax + 加权求和整体融合——FlashAttention 的核心
工具光谱(控制力递增、开发成本递增):
- torch.compile(PyTorch 2.0+):自动捕获计算图并融合,零代码改动,但融合策略保守
- Triton:Python DSL 写 block 级 kernel,接近手写 CUDA 的性能(Lecture 6 · Kernel 与 Triton 的主角)
- 手写 CUDA kernel:完全控制、性能天花板最高,但开发与维护成本大
说明
Transformer 训练中约 30-50% 的时间花在非矩阵乘法的小算子上(LayerNorm、GELU、Softmax、Dropout 等),尽管它们只占总 FLOPs 的个位数百分比——因为它们全是 memory-bound,实际耗时远超 FLOPs 占比。融合这些算子是端到端加速的关键,FlashAttention 是融合的极致案例。
5. 重计算(Recomputation / Activation Checkpointing)
5.1 内存 vs 计算的权衡
反向传播需要前向的中间结果(激活值,activations)来计算梯度——例如 $y = Wx$ 对 $W$ 的梯度需要 $x$。朴素做法是前向时把所有层的激活值存在 HBM 里,等反向时取用:
- 内存消耗:$O(\text{层数} \times \text{batch} \times \text{序列长度} \times d_{\text{model}})$,随 batch 和序列长度线性膨胀
- 大模型 + 长序列 + 大 batch 下,激活值内存可以远超参数本身(Llama 3 405B 的 BF16 权重就已 810 GB,不加检查点的完整激活值更是 TB 级)——单卡 80 GB 完全装不下
重计算(recomputation / gradient checkpointing):
- 前向时只保存少数检查点(checkpoints)(如每个 Transformer block 的输入)
- 反向传播走到某层时,从最近的检查点重新前向计算出所需激活值,用完即弃
# 伪代码(梯度检查点 / gradient checkpointing)
def forward_with_checkpointing(x, layers):
checkpoints = [x]
for i, layer in enumerate(layers):
x = layer(x)
if i % CHECKPOINT_INTERVAL == 0: # 每隔几层存一个
checkpoints.append(x)
# 否则不存,反向时重算
return x, checkpoints
def backward_with_recomputation(checkpoints, layers):
for i in range(len(layers) - 1, -1, -1):
if i not in checkpoint_indices:
# 从上一个 checkpoint 重新前向计算到这里
x = recompute_forward(checkpoints[prev], layers[prev:i+1])
grad = backward_layer(x, layers[i])
代价的定量分析:设一次前向的计算量为 $F$,反向约为 $2F$(要对输入和权重各求一次梯度),基线总计算量为 $3F$。完全重计算最多把前向再做一遍($+F$),总量变为 $4F$:
$$\frac{4F}{3F} \approx 1.33$$
即最多多花 33% 的 FLOPs,换取激活内存从 $O(L)$ 降到 $O(L/k)$ 甚至 $O(\sqrt{L})$($L$ 为层数,$k$ 为检查点间隔)。实测端到端通常只慢 10-20%——因为省下的内存允许更大的 batch,反而提高了 GPU 利用率。
提示
内存和算力是不对称的资源:内存是硬约束(80 GB 是墙,超了直接 OOM),算力是软约束(989 TFLOPS 常常用不满,尤其在 memory-bound 阶段)。用富余的算力换稀缺的内存,几乎总是划算的交易。这也解释了为什么重计算是大模型训练的标配而非无奈之举。
精细化策略:
- Selective recomputation:只重算「便宜但占内存」的算子(LayerNorm、GELU、attention softmax——都是 memory-bound、重算快),保留「昂贵」的矩阵乘法结果不重算——用最小的计算代价换最大的内存收益
- FlashAttention 本身内嵌了重计算:前向根本不存 $O(N^2)$ 的注意力矩阵,反向时按块重算(见下一节)
- 多卡场景下还可以把激活值 offload 到 CPU 内存或分片到其他卡(见 Lecture 7 · 并行 I:基础 的 ZeRO/FSDP)
6. FlashAttention:Tiling + Fusion + Recomputation 的集大成
6.1 标准 Attention 的内存瓶颈
标准缩放点积注意力(scaled dot-product attention):给定 $Q, K, V \in \mathbb{R}^{N \times d}$($N$ 为序列长度,$d$ 为每个 head 的维度):
$$S = \frac{QK^{\top}}{\sqrt{d}} \in \mathbb{R}^{N \times N}, \qquad P = \operatorname{softmax}(S) \in \mathbb{R}^{N \times N}, \qquad O = PV \in \mathbb{R}^{N \times d}$$
$S$ 是注意力分数矩阵(每对 token 之间打一个分),softmax 按行归一化得到注意力权重 $P$,最后加权求和 value 得到输出 $O$。除以 $\sqrt{d}$ 是为了让分数方差与 $d$ 无关、防止 softmax 饱和。
问题:
1. 内存 $O(N^2)$:$S$ 和 $P$ 都是 $N \times N$。$N = 4096$ 时单个 head 就是 $4096^2 \times 2\ \text{B} = 32\ \text{MB}$,乘上几十个 head 和 batch,直接爆炸;且随 $N$ 平方增长,长上下文完全不可行
2. 多次 HBM 往返:算 $S$(读 $Q,K$ 写 $S$)→ softmax(读 $S$ 写 $P$)→ 算 $O$(读 $P,V$ 写 $O$),$N \times N$ 的大矩阵在 HBM 进出多轮;softmax 每元素只做一次 exp 和除法,是典型 memory-bound 算子
6.2 FlashAttention 的三个核心技巧
FlashAttention(Dao et al., 2022) 的目标:数学上完全等价地计算 attention,但全程不物化(materialize)$O(N^2)$ 的 $S$ 和 $P$。
技巧 1:Tiling(分块计算)——把 $Q, K, V$ 沿序列维切成小块(如每块 128 行),每次只把一对块加载进 SRAM 计算。
技巧 2:Online Softmax(在线 softmax,Milakov & Gimelshein, 2018)——分块的难点在 softmax:
$$\operatorname{softmax}(x)_i = \frac{e^{x_i - m}}{\sum_j e^{x_j - m}}, \qquad m = \max_j x_j$$
分母(和减去的最大值 $m$)依赖整行数据,但分块时每次只能看到一段。解法:维护「目前为止」的最大值 $m$ 与归一化因子 $\ell$,每来一个新块就重新标定(rescale)已有的累积量。看到第 $t$ 个块 $x^{(t)}$ 时:
$$m^{(t)} = \max\left(m^{(t-1)},\ \max_j x_j^{(t)}\right)$$
$$\ell^{(t)} = \ell^{(t-1)} \, e^{m^{(t-1)} - m^{(t)}} + \sum_j e^{x_j^{(t)} - m^{(t)}}$$
其中 $e^{m^{(t-1)} - m^{(t)}}$ 是修正因子:如果新块刷新了最大值,旧的累积和是按过时的 $m$ 算的,乘这个因子把它折算到新基准下。这样单遍扫描即可得到正确的 softmax 归一化。
技巧 3:融合(Fusion)——$QK^\top$、softmax、乘 $V$ 三步在一个 kernel 内完成,$S$ 和 $P$ 的分块只短暂存在于寄存器/SRAM,输出累积器 $O$ 也随块同步 rescale:
# 伪代码(FlashAttention 前向,简化版,忽略 causal mask 等细节)
BLOCK_SIZE = 128
for i in range(0, N, BLOCK_SIZE): # 遍历 Q 的块(外层可并行)
Q_block = Q[i : i+BLOCK_SIZE, :] # (128, d) 加载到 SRAM
O_block = zeros(BLOCK_SIZE, d) # 输出累积器
m_block = -inf; l_block = 0 # 行最大值与归一化因子
for j in range(0, N, BLOCK_SIZE): # 遍历 K, V 的块
K_block = K[j : j+BLOCK_SIZE, :] # (128, d)
V_block = V[j : j+BLOCK_SIZE, :]
S_block = Q_block @ K_block.T / sqrt(d) # (128, 128) 只存在于片上
# Online softmax 更新(逐行)
m_new = max(m_block, max(S_block, dim=-1))
l_new = l_block * exp(m_block - m_new) + sum(exp(S_block - m_new), dim=-1)
# 输出累积器同步 rescale 后累加本块贡献
O_block = O_block * exp(m_block - m_new) + exp(S_block - m_new) @ V_block
m_block = m_new; l_block = l_new
O[i : i+BLOCK_SIZE, :] = O_block / l_block # 最后统一除以归一化因子,写回 HBM
反向传播怎么办? 反向也需要 $P$,但 FlashAttention 不存它——只为每行保存两个标量 $m$ 和 $\ell$(合并为 logsumexp),反向时按块重算 $S$ 和 $P$。这正是第 5 节重计算思想的内嵌应用:重算这些 memory-bound 的量比从 HBM 读 $O(N^2)$ 矩阵更快。
效果对比:
| 对比维度 | 标准 Attention | FlashAttention |
|---|---|---|
| 额外内存 | $O(N^2)$(存 $S$ 和 $P$) | $O(N)$(每行只存 $m, \ell$) |
| HBM 访问 | 大矩阵多轮往返 | $Q,K,V$ 读一次、$O$ 写一次 |
| SRAM 使用 | 基本未利用 | 显式分块调度 |
| 算术强度 | 低(softmax 部分 ~1-5 FLOPs/byte) | 高(整体接近 GEMM) |
| 实测速度 | 基线 | 2-4× 快(序列越长优势越大) |
| 训练内存 | 基线 | 5-20× 节省(取决于序列长度) |
提示
FlashAttention 没有改变 attention 的数学(结果 bit 级近似等价,不是稀疏/低秩近似),改变的是计算的编排顺序:通过 tiling + online softmax 的代数重组,把「三个 memory-bound kernel 之间用 HBM 传递 $N \times N$ 矩阵」变成「一个 kernel 内用 SRAM 传递小块」。它证明了一个重要范式:IO 感知(IO-awareness)的算法重排,可以让同一个数学函数的速度差一个数量级。
graph TB
subgraph STD["标准 Attention"]
Q1["Q K V<br/>从 HBM 读"] --> S1["算 S 矩阵<br/>写回 HBM"]
S1 --> P1["softmax 得 P<br/>写回 HBM"]
P1 --> O1["P 乘 V 得 O<br/>写回 HBM"]
O1 -.多次往返.-> HBM1["HBM<br/>3.35 TB/s"]:::slow
end
subgraph FLASH["FlashAttention"]
Q2["Q K V<br/>分块加载"] --> SRAM["SRAM 内完成<br/>S 与 P 不落 HBM"]:::fast
SRAM --> O2["O 一次写回"]
O2 --> HBM2["HBM 访问<br/>大幅减少"]:::fast
end
STD -.memory-bound.-> FLASH
classDef slow fill:#ffcdd2,stroke:#c62828
classDef fast fill:#c8e6c9,stroke:#2e7d32
6.3 FlashAttention-2/3 与后续优化
FlashAttention-2(Dao, 2023)在算法等价的前提下重排了工程实现:
- 减少非 matmul 操作:调整 rescale 的时机(延迟到最后统一做),让更高比例的时间花在 Tensor Core 上——非 matmul FLOPs 虽少,但吞吐比 Tensor Core 低一个数量级,占比稍高就拖慢整体
- 并行维度改进:除了 batch 和 head,还沿序列长度维切分给不同 thread block——长序列 + 小 batch 时也能填满所有 SM
- Warp 间任务划分:减少 warp 之间通过共享内存的中转与同步
实测比 FA-1 快 1.5-2×,比标准实现总共快 3-8×。
后续方向:
- FlashAttention-3:针对 Hopper(H100)的异步执行——用 TMA 硬件单元异步搬数据、warp 专业化(一些 warp 专职搬运、一些专职计算)重叠计算与访存,并支持 FP8
- PagedAttention(vLLM):推理时 KV cache 的分页管理,消灭显存碎片(详见 Lecture 10 · 推理)
- MQA / GQA:多个 query head 共享一组 KV head,从模型结构上减少 KV cache 的大小与带宽压力——推理时代的「结构性 FlashAttention」
7. 实战工具与性能分析
7.1 性能分析工具(Profiling)
为什么需要 profiling? 主观感觉不可靠——「这段代码应该很快」和实际快慢经常相反;瓶颈也难猜——70% 的时间可能花在你以为只占 5% 的地方。性能优化必须是测量驱动的(下一讲 Lecture 6 · Kernel 与 Triton 会完整演示这套工作流)。
| 工具 | 用途 | 输出 |
|---|---|---|
torch.cuda.Event |
测量单个操作耗时 | 毫秒级精度时间 |
torch.profiler |
端到端分析(Python + CUDA) | 每个算子的时间、内存、调用栈 |
| NVIDIA Nsight Systems | 系统级时间线(GPU + CPU) | 可视化时间线(kernel 启动、内存拷贝、空闲) |
| NVIDIA Nsight Compute | 单个 kernel 深度分析 | SM 利用率、带宽、warp 占用率、指令吞吐 |
from torch.profiler import profile, ProfilerActivity
with profile(activities=[ProfilerActivity.CPU, ProfilerActivity.CUDA]) as prof:
output = model(input)
loss = criterion(output, target)
loss.backward()
print(prof.key_averages().table(sort_by="cuda_time_total", row_limit=10))
# 输出每个算子的 GPU 时间、调用次数、内存占用
说明
- 宏观:用
torch.profiler找到最慢的几个算子 - 微观:用 Nsight Compute 分析单个 kernel 的瓶颈(compute-bound 还是 memory-bound?对照 roofline)
- 优化:针对瓶颈选择策略(融合、tiling、降精度、重计算)
- 验证:重新 profile,确认加速符合预期——优化是个闭环,不是单向操作
7.2 常见性能陷阱
| 陷阱 | 表现 | 原因 | 解决方案 |
|---|---|---|---|
| Kernel launch overhead | 大量小操作慢 | 每次 kernel 启动有 ~5-10 μs 固定 CPU 开销 | 算子融合、增大 batch、CUDA graphs |
| CPU-GPU 同步 | loss.item() 很慢 |
强制 CPU 等待 GPU 排空队列 | 减少同步点、异步 logging |
| 数据传输(PCIe) | tensor.cpu() 慢 |
PCIe 带宽 ~64 GB/s,比 HBM 慢 50 倍 | 尽量全程留在 GPU 上 |
| 未对齐/非连续内存 | 莫名其妙慢 | 非合并访问 | .contiguous()、padding、检查 strides |
| 过小的 batch size | GPU 利用率低 | 并行度不够,SM 闲置 | 增大 batch、gradient accumulation |
注意
torch.cuda.synchronize():计时必需(CUDA 异步执行,不同步测到的只是「提交任务」的时间),但留在训练代码里会拖慢训练print(tensor)、tensor.item():都会触发隐式同步- 第一次运行包含 JIT 编译与内存池初始化——benchmark 前必须 warmup(下一讲详述)
8. PyTorch 的 GPU 内存管理
8.1 内存分配策略
PyTorch CUDA caching allocator:cudaMalloc/cudaFree 是昂贵的同步操作,PyTorch 不会每次分配 tensor 都调用它们,而是先向 CUDA 要大块内存、自己在内部切分复用。释放的 tensor 内存回到 PyTorch 的池子里而不还给系统——所以 nvidia-smi 显示的占用(reserved)通常大于实际使用(allocated)。
torch.cuda.memory_allocated() # 当前 tensor 实际占用(字节)
torch.cuda.memory_reserved() # PyTorch 缓存池向 CUDA 保留的总内存
torch.cuda.max_memory_allocated() # 峰值占用(OOM 排查关键指标)
torch.cuda.reset_peak_memory_stats() # 重置峰值统计
torch.cuda.empty_cache() 只把未使用的缓存块还给系统(例如给同卡的其他进程腾地方),不会减少已分配 tensor 的内存——它救不了 OOM。
提示
Caching allocator 就是「批发转零售」:向房东(CUDA driver)一次租下整层楼(慢、要签合同/同步),再自己快速分配工位(快、纯指针操作)。代价是可能出现碎片化——池子里总空闲够但没有连续大块,这也是长序列/变长输入训练容易 OOM 的隐藏原因(可用 expandable_segments:True 缓解)。
8.2 OOM(Out of Memory)排查
典型症状:RuntimeError: CUDA out of memory
第一步:算清内存账(BF16 混合精度 + AdamW,每参数字节数):
| 组成 | 精度 | 字节/参数 |
|---|---|---|
| 权重(计算用) | BF16 | 2 |
| 梯度 | BF16 | 2 |
| 主权重副本 | FP32 | 4 |
| Adam 动量 $m$ | FP32 | 4 |
| Adam 方差 $v$ | FP32 | 4 |
| 合计(静态) | 16 |
再加上激活值(动态部分):$\text{batch} \times \text{序列长度} \times \text{层数} \times d_{\text{model}} \times \text{每激活字节数}$,通常是峰值内存的大头。
例子
静态开销 $\approx 16$ 字节/参数:
- 7B 模型:$7 \times 10^9 \times 16\ \text{B} = 112\ \text{GB}$ —— 已经超过单卡 80 GB,还没算激活值
- 单卡全参数训练的静态上限约 $80 / 16 = 5$B 参数(还要给激活值留空间)
所以「单卡训 7B」必须借助技巧:ZeRO/FSDP 把优化器状态分片到多卡(Lecture 7 · 并行 I:基础)、8-bit 优化器压缩 $m,v$、LoRA 只训少量参数、或 CPU offload。推理则只需权重 2 字节/参数 + KV cache,单卡跑 30B+ 模型毫无压力——训练与推理的内存需求相差近一个数量级。
第二步:减少内存占用(按性价比排序):
- 降低 batch size + gradient accumulation:多次 loss.backward() 累积梯度、一次 optimizer.step(),数学上等价于大 batch
- Gradient checkpointing:torch.utils.checkpoint,激活内存降 2-4×,代价 ~1.33× FLOPs(见第 5 节)
- 混合精度:确认激活值确实是 BF16 而非 FP32
- ZeRO / FSDP:把优化器状态、梯度、参数分片到多卡
第三步:可视化排查:
# 记录内存快照(PyTorch 1.13+)
torch.cuda.memory._record_memory_history()
# ... 跑到 OOM ...
torch.cuda.memory._dump_snapshot("memory_snapshot.pickle")
# 拖进 https://pytorch.org/memory_viz 可视化每块内存的分配栈
总结
mindmap
root((GPU 优化))
硬件理解
SM 与 SIMT
Warp 32 线程锁步
内存金字塔
HBM → SRAM → Register
算力 989 TFLOPS
带宽 3.35 TB/s
性能模型
Roofline 模型
算术强度 AI = FLOPs / Bytes
Compute-bound vs Memory-bound
屋脊点 ~295 FLOPs/byte
优化技巧
低精度 BF16 / FP8
内存合并 Coalescing
Tiling 分块到 SRAM
算子融合 Fusion
重计算换内存
FlashAttention
Online Softmax 单遍扫描
Tiling Q K V
SRAM 片上计算
内存从平方级降到线性级
加速 2 到 8 倍
工具链
torch.profiler
Nsight Systems / Compute
memory snapshot
gradient checkpointing
关键要点:
-
GPU = 并行机器 + 分层内存:理解 SM/SIMT/warp 执行模型、掌握 HBM ↔ SRAM ↔ Register 的延迟与带宽差异,是一切优化的前提。峰值算力 989 TFLOPS 只属于喂饱了的 Tensor Core。
-
算力 vs 带宽双重游戏:H100 算力 989 TFLOPS、带宽 3.35 TB/s,屋脊点在 AI ≈ 295 FLOPs/byte。现代训练的瓶颈往往不在矩阵乘法(compute-bound、已被 cuBLAS 优化到极限),而在小算子(LayerNorm/Softmax/GELU)的内存墙。
-
Roofline 模型指导优化方向:先算算子的算术强度、判断落在屋顶的哪一段,再针对性选择策略——memory-bound 提 AI(tiling/fusion),compute-bound 换算术单元(Tensor Core/低精度)。
-
四大优化技巧:
- 低精度(BF16/FP8):带宽减半 + Tensor Core 峰值翻倍,一举两得;BF16 保留 FP32 动态范围,免 loss scaling
- 内存合并(Coalescing):让 warp 内相邻线程访问相邻地址,事务数从 32 降到 1
- Tiling(分块):数据搬进 SRAM 反复使用,算术强度提升 tile size 倍
- 算子融合(Fusion):中间结果不落 HBM,memory-bound 算子链近似按融合数量加速 -
重计算(Recomputation):前向 $F$ + 反向 $2F$,完全重计算只加 $F$(+33% FLOPs),却把激活内存从 $O(L)$ 压到 $O(L/k)$——用富余的算力换稀缺的内存,是大模型训练的标配。
-
FlashAttention 是集大成者:Tiling(Q/K/V 分块)+ Online Softmax(增量 rescale 的单遍算法)+ Fusion($S$/$P$ 不落 HBM)+ 反向重计算(只存 logsumexp)→ 额外内存 $O(N^2) \to O(N)$、速度 2-8×,数学上完全等价。它示范了「IO 感知的算法重排」这一整个范式。
-
性能分析驱动优化:不要凭感觉猜瓶颈——
torch.profiler找最慢算子、Nsight Compute 看单 kernel 的 SM 利用率与带宽占用,测量 → 优化 → 再测量闭环推进。 -
80 GB 上限是硬约束:BF16 + AdamW 全参数训练每参数约 16 字节,单卡静态上限仅 ~5B——混合精度、gradient checkpointing、ZeRO 分片是突破单卡极限的三板斧,也是通往 Lecture 7 · 并行 I:基础 的桥梁。
下一讲预告:理解了「为什么 GPU 快」和「如何让代码跑满 GPU」,下一步是自己动手写高性能 kernel。Lecture 6 · Kernel 与 Triton 将教你:如何正确地 benchmark(warmup + synchronize)与用 torch.profiler 剖析代码、如何写 CUDA kernel(C++ 扩展)、如何用 Triton(OpenAI 开源的 Python DSL)以接近手写 CUDA 的性能编写自定义算子、torch.compile 的工作原理与适用场景,以及逐步优化 GELU 和 Softmax kernel 的实战案例。
复习自测
题目
每个输出元素:1 FLOP,搬运 6 字节(读 2 个 BF16、写 1 个),$\text{AI} = 1/6$ FLOPs/byte。代入 roofline:$P = \beta \cdot \text{AI} = 3.35\ \text{TB/s} \times \frac{1}{6} \approx 0.56$ TFLOPS,仅为峰值 989 TFLOPS 的 0.06%。结论:这类算子无论 kernel 写得多好都受带宽封顶,唯一出路是融合掉这次内存往返。
题目
位分配不同:FP16 是 5 位指数 + 10 位尾数(范围小、精度高),BF16 是 8 位指数 + 7 位尾数(范围同 FP32、精度低)。训练中溢出/下溢是灾难性的(NaN、梯度归零),而舍入误差只是 SGD 噪声的一小部分——所以「保范围弃精度」更合理。BF16 因此不需要 loss scaling,且与 FP32 互转只是截断,工程上更简单稳定。
题目
设一次前向计算量为 $F$,反向约为 $2F$(对输入、对权重各一次矩阵乘),基线总量 $3F$。完全重计算的最坏情况是反向前把前向整个重做一遍($+F$),总量 $4F$,故 $4F/3F \approx 1.33$。实际端到端只慢 10-20%,因为省下的内存换来了更大 batch、更高的 GPU 利用率。
题目
softmax 的分母 $\sum_j e^{x_j - m}$ 和最大值 $m$ 都是整行的全局量,单个块只能看到部分列,独立归一化的结果是错的。Online softmax 维护跑动的 $(m, \ell)$,每来一个新块就用修正因子 $e^{m_{\text{old}} - m_{\text{new}}}$ 把旧累积量折算到新基准,输出累积器同步 rescale——单遍扫描即得精确结果,这才使「$S$ 永不完整存在」成为可能。
题目
A.T 只改 stride 不搬数据,转置后 kernel 按「行」访问实际是跨 stride 跳读原始内存——warp 内 32 个线程的访问从连续(1 次事务)变成分散(最多 32 次事务),内存合并被破坏,有效带宽骤降。解决:.contiguous() 物化转置、或使用能处理转置布局的 GEMM 接口(cuBLAS 本身支持 op(A) 参数,此陷阱更常见于自定义 kernel 和逐元素操作)。
参考资料
- 💻 2025 Lecture 5 - GPUs.pdf(本讲讲义)
- 📄 FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness (Dao et al., 2022) — Tiling + online softmax,O(N) 内存复杂度
- 📄 FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning (Dao, 2023) — 进一步优化并行化与任务分配
- 📄 Online normalizer calculation for softmax (Milakov & Gimelshein, 2018) — Online softmax 原始论文
- 📄 Roofline: An Insightful Visual Performance Model (Williams et al., 2009) — Roofline 模型原始论文
- 📄 Mixed Precision Training (Micikevicius et al., 2018) — NVIDIA 的混合精度训练方法
- 📖 CUDA C++ Programming Guide (NVIDIA) — 官方 CUDA 文档,内存层级、SIMT 模型详解
- 📖 PyTorch CUDA Semantics 文档 — PyTorch 的 GPU 内存管理与最佳实践
- 🌐 Horace He: Making Deep Learning Go Brrrr From First Principles — compute/memory/overhead 三种瓶颈的经典科普
- 🌐 CS336 课程主页
- 🎬 FlashAttention 作者 Tri Dao 的讲解视频