# Lecture 5 · GPU

> **CS336: Language Modeling from Scratch** · Stanford · Spring 2025
> 📅 Apr 15 · 💻 [不可执行讲义（PDF）](https://github.com/stanford-cs336/spring2025-lectures/blob/main/nonexecutable/2025%20Lecture%205%20-%20GPUs.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」的核心技巧

> [!note] 一句话总览
> 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） | 大规模并行、矩阵运算 |

> [!warning] 易错点：989 TFLOPS 是有条件的峰值
> 989 TFLOPS 是 H100 SXM 的 **BF16/FP16 Tensor Core 稠密算力**。如果你的代码走的是普通 FP32 CUDA core 路径，峰值只有约 **67 TFLOPS**——差 15 倍。「用不用得上 Tensor Core」本身就是最大的优化项之一（需要低精度 dtype + 合适的矩阵形状）。

> [!tip] 直觉
> 算力差距约 **1000 倍**（989 vs ~1 TFLOPS），这就是训练语言模型必须用 GPU 的原因。但注意：**带宽只快了约 30 倍**（3.35 TB/s vs ~100 GB/s）。算力增长远快于带宽增长，意味着「喂数据」的速度越来越跟不上「算数据」的速度——这个失衡（内存墙）正是本讲一切优化技巧的根源。

```mermaid
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 为例）**：

1. **GPU 芯片（chip）**：整个物理设备
2. **流多处理器（Streaming Multiprocessor, SM）**：GPU 的「核心单元」，H100 有 **132 个 SM**（本课程实验用的 A100 是 108 个）。一个 SM 是独立的调度与执行单元，有自己的寄存器堆、共享内存和 warp 调度器
3. **CUDA 核心（CUDA core）**：每个 SM 内有 **128 个 FP32 核心**，总共 132 × 128 = **16,896 个**
4. **Tensor Core**：专用矩阵乘加单元（每周期完成一个小矩阵块的乘加），每个 SM 配备 4 个（H100 是第四代）。矩阵乘法的峰值算力几乎全部来自 Tensor Core，CUDA core 只负责其余「杂活」

**SIMT（Single Instruction, Multiple Threads）执行模型**：程序员写「单个线程做什么」，硬件负责把同一段代码复制到海量线程上并行执行。

```python
# 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 能省下控制逻辑面积的原因。

> [!warning] 易错点：分支发散（branch divergence）
> 如果同一个 warp 内的线程走了不同分支，GPU 无法「一条指令喂 32 个线程」，只能**串行执行所有分支路径**、把不满足条件的线程 mask 掉——两个分支就意味着并行度减半。
>
> ```python
> 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 的核心手艺就是显式调度共享内存。

```mermaid
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
```

> [!tip] 直觉
> 把 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、逐元素操作 |

> [!warning] 易错点：矩阵乘法快 ≠ 训练快
> 训练时间不是只由矩阵乘法决定——大量时间花在 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$——只能换更强的算术单元或降低精度

> [!example] 例：逐元素加法的性能上限
> 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 本身，这个算子都不可能更快；唯一出路是把它融合进别的算子、彻底消灭这次内存往返。

> [!tip] 直觉：四类优化策略在 roofline 上的位置
> - **算子融合（fusion）**：把多个 memory-bound 算子合并、消灭中间结果的内存往返 → 点向右移（AI 提高）
> - **Tiling（分块）**：把数据切块放进共享内存反复使用 → 点向右移
> - **低精度（BF16/FP8）**：每个数字节数减半 → AI 翻倍（点右移），同时 Tensor Core 峰值翻倍（屋顶抬高）
> - **重计算（recomputation）**：不改变单个算子的 AI，而是用富余算力换稀缺内存，让更大的 batch/模型放得下

```mermaid
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 位十进制 |

> [!tip] 直觉：为什么 BF16 赢了 FP16
> **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 下进行，避免大量小数相加的舍入误差累积

```python
# 伪代码（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**——这是它成为默认选择的重要原因。

> [!note] 实践现状
> 现代训练默认用 **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%**。

```python
# 好：合并访问（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 个元素
```

> [!warning] 易错点：转置不改变内存布局
> - 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，再在片上以任意模式读取（片上访问不受合并约束）

> [!tip] 直觉
> HBM 像超市货架：拿东西必须整箱搬（128 字节事务）。32 个人（warp）如果买的是同一箱里的 32 件商品，搬一箱就够；如果每人要的商品散布在 32 个货架上，就得搬 32 箱、每箱只用一件。合并访问就是「让结伴的线程买同一箱货」。

### 4.2 Tiling（分块）：让数据待在片上

**核心思想**：把大矩阵切成小块（tile），每块加载进共享内存后**反复使用**，把「每个数据从 HBM 读很多次」变成「读一次、片上用很多次」。

以矩阵乘法 $C = AB$（$M \times K$ 乘 $K \times N$）为例。

**朴素实现（每线程独立算一个输出）**：

```python
# 每个线程计算 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$ 维分步推进：

```python
# 伪代码（简化，忽略边界情况），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 的原因。

> [!tip] 直觉
> Tiling 的本质是**复用**：矩阵乘法天然「每个输入元素要参与 N 次计算」，朴素实现把这 N 次拆散到 N 次 HBM 读取；tiling 把它们聚拢到一次读取 + N 次片上访问。数据复用的机会是算法固有的，tiling 只是把这个机会从「L2 撞运气」变成「共享内存里明确兑现」。

```mermaid
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 传递中间结果。

```python
# 三个算子，三次 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。

```python
# 融合后：一次读 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**：完全控制、性能天花板最高，但开发与维护成本大

> [!note] 量级感受
> 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 的输入）
- 反向传播走到某层时，从最近的检查点**重新前向计算**出所需激活值，用完即弃

```python
# 伪代码（梯度检查点 / 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 利用率。

> [!tip] 直觉
> 内存和算力是不对称的资源：**内存是硬约束**（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：

```python
# 伪代码（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× 节省**（取决于序列长度）|

> [!tip] 直觉
> FlashAttention 没有改变 attention 的数学（结果 bit 级近似等价，不是稀疏/低秩近似），改变的是**计算的编排顺序**：通过 tiling + online softmax 的代数重组，把「三个 memory-bound kernel 之间用 HBM 传递 $N \times N$ 矩阵」变成「一个 kernel 内用 SRAM 传递小块」。它证明了一个重要范式：**IO 感知（IO-awareness）的算法重排，可以让同一个数学函数的速度差一个数量级**。

```mermaid
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 占用率、指令吞吐 |

```python
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 时间、调用次数、内存占用
```

> [!note] Profiling 工作流
> 1. **宏观**：用 `torch.profiler` 找到最慢的几个算子
> 2. **微观**：用 Nsight Compute 分析单个 kernel 的瓶颈（compute-bound 还是 memory-bound？对照 roofline）
> 3. **优化**：针对瓶颈选择策略（融合、tiling、降精度、重计算）
> 4. **验证**：重新 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 |

> [!warning] 易错点：测量本身会改变性能
> - `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）。

```python
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。

> [!tip] 直觉
> 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{每激活字节数}$，通常是峰值内存的大头。

> [!example] 算一算：80 GB 能全参数训练多大的模型？
> 静态开销 $\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**：把优化器状态、梯度、参数分片到多卡

**第三步：可视化排查**：

```python
# 记录内存快照（PyTorch 1.13+）
torch.cuda.memory._record_memory_history()
# ... 跑到 OOM ...
torch.cuda.memory._dump_snapshot("memory_snapshot.pickle")
# 拖进 https://pytorch.org/memory_viz 可视化每块内存的分配栈
```

---

## 总结

```mermaid
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
```

**关键要点**：

1. **GPU = 并行机器 + 分层内存**：理解 SM/SIMT/warp 执行模型、掌握 HBM ↔ SRAM ↔ Register 的延迟与带宽差异，是一切优化的前提。峰值算力 989 TFLOPS 只属于喂饱了的 Tensor Core。

2. **算力 vs 带宽双重游戏**：H100 算力 989 TFLOPS、带宽 3.35 TB/s，屋脊点在 AI ≈ 295 FLOPs/byte。现代训练的瓶颈往往不在矩阵乘法（compute-bound、已被 cuBLAS 优化到极限），而在小算子（LayerNorm/Softmax/GELU）的内存墙。

3. **Roofline 模型指导优化方向**：先算算子的算术强度、判断落在屋顶的哪一段，再针对性选择策略——memory-bound 提 AI（tiling/fusion），compute-bound 换算术单元（Tensor Core/低精度）。

4. **四大优化技巧**：
   - **低精度（BF16/FP8）**：带宽减半 + Tensor Core 峰值翻倍，一举两得；BF16 保留 FP32 动态范围，免 loss scaling
   - **内存合并（Coalescing）**：让 warp 内相邻线程访问相邻地址，事务数从 32 降到 1
   - **Tiling（分块）**：数据搬进 SRAM 反复使用，算术强度提升 tile size 倍
   - **算子融合（Fusion）**：中间结果不落 HBM，memory-bound 算子链近似按融合数量加速

5. **重计算（Recomputation）**：前向 $F$ + 反向 $2F$，完全重计算只加 $F$（+33% FLOPs），却把激活内存从 $O(L)$ 压到 $O(L/k)$——用富余的算力换稀缺的内存，是大模型训练的标配。

6. **FlashAttention 是集大成者**：Tiling（Q/K/V 分块）+ Online Softmax（增量 rescale 的单遍算法）+ Fusion（$S$/$P$ 不落 HBM）+ 反向重计算（只存 logsumexp）→ 额外内存 $O(N^2) \to O(N)$、速度 2-8×，数学上完全等价。它示范了「IO 感知的算法重排」这一整个范式。

7. **性能分析驱动优化**：不要凭感觉猜瓶颈——`torch.profiler` 找最慢算子、Nsight Compute 看单 kernel 的 SM 利用率与带宽占用，测量 → 优化 → 再测量闭环推进。

8. **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 的实战案例。

---

## 复习自测

> [!question]- Q1：算一算——BF16 逐元素加法 $C = A + B$ 在 H100 上的性能上限是多少 TFLOPS？占峰值算力的百分之几？
> 每个输出元素：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 写得多好都受带宽封顶，唯一出路是融合掉这次内存往返。

> [!question]- Q2：概念辨析——BF16 和 FP16 同为 16 位，为什么现代训练几乎都选 BF16？
> 位分配不同：FP16 是 5 位指数 + 10 位尾数（范围小、精度高），BF16 是 8 位指数 + 7 位尾数（范围同 FP32、精度低）。训练中溢出/下溢是灾难性的（NaN、梯度归零），而舍入误差只是 SGD 噪声的一小部分——所以「保范围弃精度」更合理。BF16 因此不需要 loss scaling，且与 FP32 互转只是截断，工程上更简单稳定。

> [!question]- Q3：推导——为什么完全的 activation checkpointing 只增加约 33% 的计算量？
> 设一次前向计算量为 $F$，反向约为 $2F$（对输入、对权重各一次矩阵乘），基线总量 $3F$。完全重计算的最坏情况是反向前把前向整个重做一遍（$+F$），总量 $4F$，故 $4F/3F \approx 1.33$。实际端到端只慢 10-20%，因为省下的内存换来了更大 batch、更高的 GPU 利用率。

> [!question]- Q4：为什么 FlashAttention 分块计算时必须用 online softmax，而不能各块独立算 softmax 再拼起来？
> softmax 的分母 $\sum_j e^{x_j - m}$ 和最大值 $m$ 都是**整行**的全局量，单个块只能看到部分列，独立归一化的结果是错的。Online softmax 维护跑动的 $(m, \ell)$，每来一个新块就用修正因子 $e^{m_{\text{old}} - m_{\text{new}}}$ 把旧累积量折算到新基准，输出累积器同步 rescale——单遍扫描即得精确结果，这才使「$S$ 永不完整存在」成为可能。

> [!question]- Q5：`x = A.T @ B` 突然比 `x = A @ B` 慢很多（A 形状对称），最可能的硬件层面原因是什么？
> `A.T` 只改 stride 不搬数据，转置后 kernel 按「行」访问实际是跨 stride 跳读原始内存——warp 内 32 个线程的访问从连续（1 次事务）变成分散（最多 32 次事务），内存合并被破坏，有效带宽骤降。解决：`.contiguous()` 物化转置、或使用能处理转置布局的 GEMM 接口（cuBLAS 本身支持 op(A) 参数，此陷阱更常见于自定义 kernel 和逐元素操作）。

---

## 参考资料

- 💻 [2025 Lecture 5 - GPUs.pdf（本讲讲义）](https://github.com/stanford-cs336/spring2025-lectures/blob/main/nonexecutable/2025%20Lecture%205%20-%20GPUs.pdf)
- 📄 [FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness (Dao et al., 2022)](https://arxiv.org/abs/2205.14135) — Tiling + online softmax，O(N) 内存复杂度
- 📄 [FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning (Dao, 2023)](https://arxiv.org/abs/2307.08691) — 进一步优化并行化与任务分配
- 📄 [Online normalizer calculation for softmax (Milakov & Gimelshein, 2018)](https://arxiv.org/abs/1805.02867) — Online softmax 原始论文
- 📄 [Roofline: An Insightful Visual Performance Model (Williams et al., 2009)](https://www2.eecs.berkeley.edu/Pubs/TechRpts/2008/EECS-2008-134.pdf) — Roofline 模型原始论文
- 📄 [Mixed Precision Training (Micikevicius et al., 2018)](https://arxiv.org/abs/1710.03740) — NVIDIA 的混合精度训练方法
- 📖 [CUDA C++ Programming Guide (NVIDIA)](https://docs.nvidia.com/cuda/cuda-c-programming-guide/) — 官方 CUDA 文档，内存层级、SIMT 模型详解
- 📖 [PyTorch CUDA Semantics 文档](https://pytorch.org/docs/stable/notes/cuda.html) — PyTorch 的 GPU 内存管理与最佳实践
- 🌐 [Horace He: Making Deep Learning Go Brrrr From First Principles](https://horace.io/brrr_intro.html) — compute/memory/overhead 三种瓶颈的经典科普
- 🌐 [CS336 课程主页](https://stanford-cs336.github.io/spring2025/)
- 🎬 [FlashAttention 作者 Tri Dao 的讲解视频](https://www.youtube.com/watch?v=gMOAud7hZg4)
