斯坦福 CS336 从零构建大模型 (2025 春) - 第六讲:GPU 高性能编程与 Kernel 融合

斯坦福 CS336 第六讲的核心主题是**“如何为 GPU 编写高性能代码”**。课程深入探讨了 GPU 的底层执行逻辑、代码基准测试与性能分析(Profiling)的最佳实践,并通过实现自定义的 C++ CUDA Kernel、Triton 以及使用 torch.compile,详细讲解了算子融合(Kernel Fusion)的过程。

以下是本节课所有核心知识点的全景梳理:

一、 GPU 硬件与执行模型回顾 (GPU Hardware & Execution Model)

在编写高性能代码前,讲师首先复习了 GPU 的核心架构机制:

  • 物理与逻辑层级: GPU 包含多个流式多处理器(SM),每个 SM 包含大量计算单元。执行任务时,任务被划分为多个 线程块(Thread Blocks),每个 Block 被调度到一个独立的 SM 上运行。Block 内部包含大量 线程(Threads) 来执行具体计算。
  • 为什么需要线程块 (Thread Blocks)? 线程块内的线程可以通过 SM 内部极速的**共享内存(Shared Memory)**互相通信和同步。但不同的线程块之间无法同步,只能通过极慢的全局显存(DRAM)交流。
  • 波次 (Waves/Warps): 线程是以 32 个为一组(即 Warp)作为一个波次(Wave)同时执行的。为了最大化利用率,应当确保 Block 的数量能被 SM 的数量整除,以避免部分 SM 在最后一波次中闲置。
  • 算术强度 (Arithmetic Intensity): 优化代码的核心目标是保持高算术强度(即每次内存读写伴随尽可能多的浮点运算 FLOPs)。矩阵乘法受限于算力(Compute-bound),而其他绝大多数操作都受限于内存带宽(Memory-bound)。

二、 基准测试的最佳实践 (Benchmarking)

如果想知道代码跑得多快,不能简单地用 Python 的 time 模块随便测。讲师强调了两个写测试代码时必须遵守的铁律:

  • 预热迭代 (Warm-up iterations): 第一次调用 PyTorch 的 GPU 算子时,系统在后台需要进行机器码编译、指令下发等初始化工作。必须先运行几次预热代码,才能测量到稳定状态下的真实速度。
  • 强行同步 torch.cuda.synchronize(): CPU 和 GPU 是异步工作的。Python 代码在 CPU 上运行,它只会把指令“排队”发送给 GPU 然后继续往下跑,不会等待 GPU 算完。如果不加同步锁,你测出的极短时间只是“CPU 发送指令的时间”,而非“GPU 实际计算的时间”。调用同步指令可以强制 CPU 停下,直到 GPU 执行完毕,从而获得准确耗时。

三、 性能分析与 CPU/GPU 异步陷阱 (Profiling & The Asynchrony Trap)

基准测试只能告诉你“代码有多慢”,而**性能分析器(Profiler)**能告诉你“时间花在了哪里”。

  • PyTorch 内置 Profiler:
    • 通过它可以看到 Python 层面的操作(如 a + b)是如何被转换成 C++ 接口(A10),最终派发给特定的底层 CUDA Kernel(如 cutlass 矩阵乘法库)的。
    • 根据矩阵尺寸的不同,PyTorch 会自动分发给不同的底层算子(例如小矩阵可能会使用 xmma_gmm 而不是 cutlass)。
  • NSight Systems (NSys) 与隐藏的性能杀手:
    • 使用 NVIDIA 的专业分析器可以看到极细粒度的 CPU 和 GPU 时间轴。
    • Print 语句陷阱: 因为 CPU 执行速度远快于 GPU 计算,正常情况下 CPU 会超前把很多步骤的指令排入 GPU 队列。但如果你在训练循环中加了一句 print(loss),CPU 就必须停下来干等 GPU 算完 Loss 传回来才能打印。这会彻底打破 CPU/GPU 的异步流水线,导致 GPU 出现等待指令的闲置期(即制造了人为的 CPU 瓶颈)。

四、 算子融合与手写底层 Kernel (Kernel Fusion)

为了证明算子融合(Kernel Fusion,即将多个小计算合并成一个大的 CUDA 任务以减少显存读写)的威力,讲师以 GLU(门控线性单元) 激活函数为例进行了实战演练:

  • Naive PyTorch 实现 (最慢): 如果直接用 Python 组合乘法、加法、tanh 和幂运算,系统会启动多个独立的 CUDA Kernel。这导致数据在显存和 SM 之间被反复来回搬运,耗时高达 8.1 毫秒。
  • PyTorch 官方融合版 (最快): 直接调用 torch.nn.functional.gelu,只触发一次底层的专用融合 Kernel,耗时仅 1.1 毫秒。
  • 手写 C++ CUDA Kernel:
    • 为了学习原理,课程展示了如何在 C++ 中写一个融合 Kernel。
    • 核心逻辑: 检查显存是否连续(is_contiguous);计算需要划分多少个 Grid 和 Block(总量除以 Block 尺寸向上取整);在 Kernel 内部,通过 blockIdx.x * blockDim.x + threadIdx.x 算出当前线程应该处理哪个具体元素,并进行边界检查防止越界溢出。
    • 结果:将速度提升到了约 1.8 毫秒。
  • 使用 Triton 编写 Kernel:
    • Triton 是 OpenAI 开发的 Python 级 DSL。它屏蔽了最底层的线程(Thread)管理,让开发者在“Block(块)”的层面上写代码。
    • 优势: 全 Python 编写,自动处理内存合并访问(Memory coalescing,即一次性拉取相邻的4个字节以触发突发模式读取)和共享内存管理。
    • PTX 分析: 讲师带读了 Triton 编译后的底层 PTX 机器码,证实了它确实在巧妙地使用极速的寄存器(Registers)暂存数据,并成块地读取内存。
    • 结果:耗时也是 1.8 毫秒左右,但代码编写难度极大幅度下降。
  • 终极武器:torch.compile:
    • 这是 PyTorch 现代的 JIT 编译器。它能在底层自动帮你把那些 Naive 的 Python 零碎操作融合成高效的 Triton 代码。
    • 结果:自动优化后的耗时降至 1.47 毫秒,甚至比手写的基础 Triton 还要快。
    • 教训: 除非是像 Flash Attention 这样极度复杂的架构级创新(torch.compile 无法自动推导出正确的硬件级复用策略),否则日常绝大部分的算子融合工作,直接交给 torch.compile 就足够了。

五、 Triton Softmax 实战示例 (Triton Softmax)

在课程最后,讲师展示了如何用 Triton 写一个比 GLU 稍微复杂的算子——Softmax。

  • 难点: Softmax 不是逐个元素独立操作的,它需要对整行数据进行归约求和(Reduction)。
  • Naive Triton 设计: 最简单高效的方法是让 Grid = Rows,即把矩阵的每一行分配给一个独立的 Block(SM) 处理。将 Block 的大小设为大于等于列数的下一个 2 的幂。这样就可以一次性把整行加载到 SM 极速的共享内存中,在内部完成求最大值、指数运算、求和与归一化除法,最后再写回全局显存,从而完美解决内存瓶颈

六、核心概念问答 (Q&A)

Q1:为什么 print(loss) 会打断 CPU/GPU 异步流水线并导致 GPU 闲置?

正常情况:CPU 和 GPU 异步流水线

CPU 和 GPU 是两个独立的处理器,可以同时工作:

CPU: [发指令1] [发指令2] [发指令3] [发指令4] …
GPU: [执行1] [执行2] [执行3] …

CPU 不等 GPU 执行完,就继续发下一条指令。
GPU 有一个指令队列,慢慢消化。

这就像:

老板(CPU):不断把任务单贴到公告板上
工人(GPU):按顺序取任务单慢慢干

老板不用等工人干完,继续贴下一张任务单。


print(loss) 发生了什么

loss 是一个存在 GPU 显存里的数值。

print 是 CPU 的操作,要打印必须先把数据从 GPU 传到 CPU。

所以:

CPU:[发指令1] [发指令2] … [要打印了!] ← 停!
等 GPU 算完 loss
等数据传过来
打印
[继续发指令]

GPU:[执行1] [执行2] … [算完loss] [等新指令] ← 闲置!


为什么会导致 GPU 闲置

正常:CPU 总是超前,GPU 队列里永远有任务

print 之后:
CPU 被迫等 GPU
→ CPU 停止发新指令
→ GPU 队列被清空
→ GPU 算完 loss 之后无事可做
→ GPU 闲置,浪费算力


正确做法

不要每步都 print,而是每 N 步打印一次,或者用异步日志:

慢(每步都同步)
for step in range(10000):
loss = train_step()
print(loss) # 每步都打断流水线
快(每100步打印)
for step in range(10000):
loss = train_step()
if step % 100 == 0:
print(loss) # 99步不打断,影响小很多


一句话总结

print(loss) 强迫 CPU 等 GPU 把结果传回来,打断了 CPU 超前发指令的节奏,导致
GPU 队列清空后无事可做,形成闲置。


七、第六讲复习题 (Lecture 6: Kernels, Triton)

一、 基准测试与性能分析 (Benchmarking & Profiling)

  1. 预热 (Warm-up): 在测试 PyTorch GPU 代码的运行时间时,为什么必须在开始计时前执行几次“预热”迭代?
  2. 异步陷阱 (The Asynchrony Trap): 如果在测量 GPU 函数执行时间时,忘记调用 torch.cuda.synchronize(),你会得到什么错误的测量结果?为什么?
  3. Print 语句的致命瓶颈: 在正常的深度学习训练循环中,CPU 通常会超前 GPU 运行(将指令排入队列)。如果在训练循环中加入一句简单的 print(loss),会对系统性能造成什么毁灭性的打击?
  4. 底层算子分发: 性能分析器(Profiler)显示,在 PyTorch 中执行简单的矩阵乘法(A @ B)时,底层并不会总是调用同一个 CUDA Kernel。系统(例如使用 torch.compile 时)会根据什么条件来动态决定分发给哪一个特定的底层算子(如 cutlass 或 xmma_gmm)?

二、 算子融合与底层原理 (Kernel Fusion & Under the Hood)

  1. 算子融合的威力: 如果用纯 Python(乘法、加法、tanh 拼接)写一个 GLU 激活函数,耗时约 8.1 毫秒;而调用官方的融合算子只需 1.1 毫秒。请从 GPU 硬件的角度解释,为什么纯 Python 拼接写法会如此之慢?
  2. 内存连续性 (Contiguous Memory): 在将 PyTorch Tensor 传入自定义的 C++ CUDA Kernel 或 Triton Kernel 之前,为什么通常必须检查或强制要求该 Tensor 是连续的(is_contiguous)?

三、 手写 CUDA C++ (Writing CUDA in C++)

  1. 线程坐标计算: 在编写基础的 1D CUDA Kernel 时,GPU 不会直接告诉当前线程它负责处理数组中的第几个元素。程序员必须使用三个内置变量来计算全局索引 i,请写出这个标准的索引计算公式。
  2. 越界检查 (Bounds Checking): 在 CUDA Kernel 内部计算出全局索引 i 后,必须立刻进行什么极其重要的条件判断?如果不做会导致什么后果?

四、 Triton 与现代编译器 (Triton & torch.compile)

  1. Triton 的抽象层级: 原生的 CUDA 要求开发者精细管理每一个线程 (Thread)。而 OpenAI 开发的 Triton 语言提升了抽象层级,在 Triton 中,开发者编写逻辑的基本操作单元(抽象单位)是什么?
  2. Triton 实现 Softmax: Softmax 需要计算一整行的最大值和总和(Reduction 操作)。在用 Triton 编写高效的 Softmax Kernel 时,为了让这一行的规约计算全部在极速的共享内存中完成,通常会如何映射“矩阵的行”与 Triton 的执行单元?

八、参考答案与知识点解析

  1. 为什么需要预热 (Warm-up)?

    答案: 第一次在 PyTorch 中调用 GPU 操作时,系统在后台需要进行大量的初始化工作,包括机器码的即时编译(JIT)、环境加载以及将指令发送到 GPU。如果不预热,你测量到的时间将包含这些极其耗时的初始化开销,而无法反映代码在稳定状态下的真实运行速度。

  2. 忘记 torch.cuda.synchronize() 的后果?

    答案: 你测出的时间会异常的短(并且是错误的)。因为 CPU 和 GPU 是异步工作的。如果没有同步锁,CPU 只会极快地把计算指令“扔进” GPU 的执行队列中,然后立刻返回并结束计时,此时 GPU 甚至可能还没开始计算。调用 synchronize() 会强制 CPU 停下等待,直到 GPU 执行完所有队列中的任务,从而测出真实的计算时间。

  3. print(loss) 造成的性能瓶颈?

    答案: 正常情况下,CPU 执行速度极快,会超前把大量 Kernel 调度指令压入 GPU 队列,让 GPU 保持满载。但是,print(loss) 发生在 CPU 上,且它需要 GPU 算出的具体数值。这会强制系统插入一个隐式的同步点(Synchronization barrier),CPU 必须停下来干等 GPU 把当前的 Loss 算完并传回内存。这彻底打破了 CPU 的超前调度流水线,导致 GPU 出现大量闲置等待时间,制造了人为的 CPU 瓶颈。

  4. 底层算子分发的依据?

    答案: 底层分发主要取决于矩阵的维度和大小。例如,对于巨大的矩阵,PyTorch 可能会调用 NVIDIA 的 cutlass 库;而对于 128x128 这样较小的矩阵,则会分发给特定的 xmma_gmm 算子。不同的尺寸和硬件特征对应着不同的最优切块(Tiling)和调度策略。

  5. Python 拼接写法极慢的硬件原因(算子融合的本质)?

    答案: 纯 Python 的拼接写法(如先算加法,再算乘法,再算幂)会启动多个独立的 CUDA Kernel。每执行完一步,GPU 都必须将庞大的中间结果写回极慢的全局显存(DRAM),下一步再读出来。算子融合(Kernel Fusion)将所有操作打包进一个 Kernel,在计算单元极速的局部缓存(SRAM/寄存器)中一次性算完所有步骤,只进行一次全局显存读写,极大地消除了内存带宽瓶颈。

  6. 为什么要求内存连续 (Contiguous)?

    答案: CUDA Kernel 底层依赖于一维的、线性的指针算术(Pointer arithmetic)来访问数据。如果 Tensor 被转置过或被切片(View/Slice),它在逻辑上虽然是矩阵,但在物理显存中的地址却是不连续的(有步长跳转/Stride)。如果 Kernel 按照连续的平铺逻辑去强行读取,就会取到错误的数据。

  7. 线程坐标计算公式?

    答案: i = blockIdx.x * blockDim.x + threadIdx.x。 (解析:blockIdx.x * blockDim.x 算出当前线程块在全局中的起始偏移量,加上 threadIdx.x 这个块内偏移,就得到了该线程负责处理的全局元素的绝对索引。)

  8. 越界检查 (Bounds Checking)?

    答案: 必须立刻判断 if (i < num_elements)。 (解析:因为线程块的尺寸通常是固定的(如 128 或 256 的倍数),总元素数量往往不能被完美整除,系统会向上取整分配多余的线程。最后一个块中的尾部线程如果执行操作,就会读写越界(访问未分配的内存),导致程序崩溃或产生垃圾数据。)

  9. Triton 的抽象层级?

    答案: Triton 将开发者编程的抽象层级提升到了**线程块(Thread Blocks / Blocks)**级别。开发者直接编写处理一整个 Block 数据的向量化操作(Vectorized operations),而底层的共享内存分配、线程同步以及显存合并访问(Memory coalescing)等极其繁琐的细节,全部由 Triton 编译器在后台自动优化和管理。

  10. Triton 实现 Softmax 的映射策略?

    答案: 通常的策略是让每一个块(Block)精确负责处理矩阵中的一行(Row)。因为 Softmax 需要对一整行求最大值、求和并相除,让一整行作为一个 Block 加载到同一个 SM(流式多处理器)中,就可以直接在极速的局部共享内存里完成所有规约计算,最后再统一写回全局显存,大幅提升效率。

更多推荐