CS336 LecTure 6 笔记补充

GPU 硬件架构与执行模型

要写出高性能代码,必须理解 A100/H100 等 GPU 硬件结构:

  1. 流式多处理器 (SM) 与核心:GPU 由数十上百个 SM 组成(如 A100 有 128 个 SM)。每个 SM 包含大量的 FP32/FP16/BF16 计算核心(CUDA Core / Tensor Core)。
  2. 金字塔式分层存储架构
    • DRAM (Global Memory):容量极大(如 80GB),但带宽受限(A100 约 2TB/s)。所有跨线程块(Block)的数据交换必须经过这里。
    • SRAM (Shared Memory / L1 Cache):每个 SM 独有,容量极小(A100 单个 SM 有 192KB),但速度极快(大概 20TB/s)。用于 Block 内线程间的数据共享与通信。
    • Registers (寄存器):每个线程私有,速度最快(1个时钟周期就能访问)。高性能 Kernel 的终极目标就是最大限度的把中间变量驻留在寄存器里。
  3. 线程层级与 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\) 的高精度近似公式:

\[\text{GELU}(x) \approx 0.5x \left( 1 + \tanh\left[ \sqrt{\frac{2}{\pi}} \left( x + 0.044715x^3 \right) \right] \right) \]

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-computensight-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 读写,然后提升性能。

posted @ 2026-07-23 08:46  PassName  阅读(2)  评论(0)    收藏  举报