CS336 LecTure 6 笔记补充
GPU 硬件架构与执行模型
要写出高性能代码,必须理解 A100/H100 等 GPU 硬件结构:
- 流式多处理器 (SM) 与核心:GPU 由数十上百个 SM 组成(如 A100 有 128 个 SM)。每个 SM 包含大量的 FP32/FP16/BF16 计算核心(CUDA Core / Tensor Core)。
- 金字塔式分层存储架构:
- DRAM (Global Memory):容量极大(如 80GB),但带宽受限(A100 约 2TB/s)。所有跨线程块(Block)的数据交换必须经过这里。
- SRAM (Shared Memory / L1 Cache):每个 SM 独有,容量极小(A100 单个 SM 有 192KB),但速度极快(大概 20TB/s)。用于 Block 内线程间的数据共享与通信。
- Registers (寄存器):每个线程私有,速度最快(1个时钟周期就能访问)。高性能 Kernel 的终极目标就是最大限度的把中间变量驻留在寄存器里。
- 线程层级与 SIMT 架构:
- Grid -> Block -> Thread:GPU 调度以线程块(Block)为单位,一个 Block 会被分配到一个 SM 上执行。
- Warp:SM 调度的最小硬件单位。每 32 个连续线程组成一个 Warp。Warp 内的 32 个线程遵循 SIMT(单指令多线程) 模型,在同一时钟周期执行完全相同的指令。
- 分支发散 (Branch Divergence):如果 Warp 内的线程走了不同的
if-else路径,GPU 必须串行执行每个分支(先屏蔽一部分线程执行分支A,再屏蔽另一部分执行分支B),这将导致严重的性能下降。上一节课里我们提到了,每个线程走的一定是同一个程序,所以尽量不要出现if-else
GELU
1. GELU 的数学本质与工程近似
传统的 ReLU 是硬截断(\(x \cdot \mathbb{I}(x>0)\)),会导致死神经元。GELU 引入了概率化软门控:\(\text{GELU}(x) = x \cdot \Phi(x)\),其中 \(\Phi(x)\) 是标准正态分布的累积分布函数。
精确公式误差函数积分:\(\text{GELU}(x) = \frac{x}{2} [ 1 + \text{erf}( \frac{x}{\sqrt{2}} ) ]\)。
在 GPU 硬件上,直接计算 erf 不太现实。工程上普遍采用基于 \(\tanh\) 的高精度近似公式:
2. Triton 手写 GELU Kernel
Triton 允许用类 Python 语法直接操作 Block 级别的数据,编译器会自动处理底层的 Warp 调度和 Shared Memory
import triton
import triton.language as tl
@triton.jit
def gelu_kernel(
x_ptr, # 输入张量在 DRAM 中的指针
output_ptr, # 输出张量在 DRAM 中的指针
n_elements, # 张量总元素个数
BLOCK_SIZE: tl.constexpr, # 编译期常量:当前 Block 负责处理的元素数(如 1024)
):
# 计算当前程序 ID(对应一个 Block)在网格中的位置
pid = tl.program_id(axis=0)
# 计算当前 Block 负责的数据在 DRAM 中的起始偏移量
block_start = pid * BLOCK_SIZE
# 生成当前 Block 内所有线程需要处理的全局内存偏移索引
# 例如生成 [0, 1, ..., BLOCK_SIZE-1] 并加上起始偏移
offsets = block_start + tl.arange(0, BLOCK_SIZE)
# 边界检查掩码(Mask):防止处理最后一个 Block 时发生数组越界
mask = offsets < n_elements
# 从 DRAM 加载数据
# 使用 mask 进行保护加载。数据被一次性读入寄存器(Register)
x = tl.load(x_ptr + offsets, mask=mask)
# 关键步骤:寄存器内计算
# 以下所有运算(乘、加、tanh)全部在 SM 的寄存器中完成,不访问 DRAM!
# 映射公式:0.5 * x * (1 + tanh(sqrt(2/pi) * (x + 0.044715 * x^3)))
cdf = 0.5 * (1.0 + tl.math.tanh(0.7978845608 * (x + 0.044715 * x * x * x)))
output = x * cdf
# 关键步骤:写回 DRAM
# 将寄存器中的最终计算结果,通过 mask 保护,合并写回全局显存
tl.store(output_ptr + offsets, output, mask=mask)
性能分析 Profiling 与基准测试 Benchmarking
写完 Kernel 后,必须通过科学的方法验证其性能是否达到了硬件极限。
1. 正确的 Benchmarking 姿势
在 PyTorch/Triton 中测量 GPU 耗时,必须注意的两个问题:
- 异步陷阱:CPU 发出 GPU 指令后不会等待,直接继续执行。必须使用
torch.cuda.synchronize()或 CUDA Events 来确保 GPU 计算完成后再记录结束时间。 - 冷启动(JIT编译)陷阱:首次运行 Kernel 时,GPU 驱动和 Triton/PyTorch 编译器需要进行编译和初始化。所以应该先进行 5-10 次 Warmup(热身迭代),再进行正式计时。
2. Profiling 性能分析核心指标
使用 nsight-compute 或 nsight-systems 等工具分析 Kernel 时,重点关注:
- Occupancy (占用率):SM 上活跃的 Warp 数量占硬件理论最大值的比例。
- L1/Shared Memory Throughput:检查是否充分利用了片上高速缓存。
- DRAM Throughput (内存带宽利用率):对于 GELU、Softmax 这类算子,如果带宽利用率达到了硬件上限(如 A100 接近 2TB/s),说明 Kernel 已经优化到极致(Memory-bound 算子的终极目标)。
自动化优化:PyTorch 2.0 torch.compile
除了手写 Triton/CUDA,课程最后还讲了一嘴 PyTorch 的即时编译 JIT 方案。
当你调用 model = torch.compile(model) 时,PyTorch 2.0 背后的 TorchDynamo 会捕获计算图,并通过后端编译器(如 Inductor,底层经常生成 Triton 代码)自动执行算子融合。它能自动识别出 matmul -> bias -> gelu ,并编译成一个单独的融合 Kernel,在无需修改业务代码的情况下,减少 DRAM 读写,然后提升性能。

浙公网安备 33010602011771号