工作台课程

CS336 · 从零构建语言模型

Lecture 5 · GPU

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 为例)

  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)执行模型:程序员写「单个线程做什么」,硬件负责把同一段代码复制到海量线程上并行执行。

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

说明

  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

注意

  • torch.cuda.synchronize():计时必需(CUDA 异步执行,不同步测到的只是「提交任务」的时间),但留在训练代码里会拖慢训练
  • print(tensor)tensor.item():都会触发隐式同步
  • 第一次运行包含 JIT 编译与内存池初始化——benchmark 前必须 warmup(下一讲详述)

8. PyTorch 的 GPU 内存管理

8.1 内存分配策略

PyTorch CUDA caching allocatorcudaMalloc/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 checkpointingtorch.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

关键要点

  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」,下一步是自己动手写高性能 kernelLecture 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 和逐元素操作)。


参考资料