/ CS336
[CS336-06] Kernels 与 Triton
整理 CS336 第六讲:benchmark、profiling、GPU execution model、Triton kernel、tiling 和 kernel fusion。
这篇是 Stanford CS336 Spring 2025 第 6 讲 Course Materials 的学习笔记。主题是 benchmark、profile 和手写 kernel。
1. 从 GPU 原理到 kernel 实现#
第 5 讲建立了 GPU 的 memory hierarchy 和 execution model,第 6 讲把这些原则具体落实到代码:如何测一个操作的速度,如何判断是 compute-bound 还是 memory-bound,如何用 Triton 写 kernel。
核心原则仍然是:
组织计算,最小化 global memory 的读写。2. Benchmark 要测对东西#
benchmark 一个 kernel 时,不能只看一次运行时间。GPU 执行是异步的,计时需要同步;第一次运行可能包含编译和 warmup;输入 shape、dtype、stride、contiguity 都会显著影响结果。
因此一个可靠 benchmark 至少要考虑:
- warmup 和重复测量;
- CUDA synchronize;
- 统计均值、方差或 percentile;
- 数据规模是否足够大;
- 是否包含不该算进去的 CPU overhead。
这也是为什么系统优化中 profile 比直觉重要。很多 bottleneck 只有在真实 trace 里才明显。
3. Execution model:blocks、warps、occupancy#
GPU kernel 被拆成 thread blocks,blocks 被调度到 SM 上执行。一个 block 内的线程共享 shared memory,warp 是实际执行调度的重要单位。
如果 grid/block 设置不合理,就可能出现 occupancy 低、last wave 不满、warp divergence、memory access 不连续等问题。
这里的难点是:编程模型暴露了一部分硬件信息,但隐藏了很多调度细节。因此 kernel 性能往往需要实验、profile 和反复调整。
4. Tiling:把数据搬进 shared memory 复用#
矩阵乘法是最重要的例子。朴素实现会反复从 global memory 读取同一批数据。Tiling 的思想是把矩阵拆成小块,让一个 block 负责一个 output tile,并把相关 A/B tile 放进 shared memory 中复用。
这样可以显著提高 arithmetic intensity:同样一次 global memory 读取,被更多 FLOPs 使用。
Triton 的价值是让我们用接近 Python 的方式表达 block-level program,同时仍能控制 program id、mask、load/store、tile shape 等底层细节。
5. Kernel fusion#
很多 Transformer 操作不是大 matmul,而是 RMSNorm、activation、elementwise add、mask、softmax 等小操作。如果每个操作都是独立 kernel,就会不断把中间结果写回 global memory,再读出来。
Kernel fusion 把多个操作合并进一个 kernel,减少中间 memory traffic。Assignment 2 里实现 fused RMSNorm,就是这个思想的练习。
6. takeaway#
第 6 讲把 GPU 优化从概念带到实践:
- benchmark 要严格,否则数字容易骗人;
- profile 是定位瓶颈的必要工具;
- kernel 性能来自 memory movement、occupancy、tiling、fusion 的综合结果;
- Triton 提供了介于 PyTorch 和 CUDA 之间的可控抽象。
对大模型工程来说,会写 PyTorch 只能保证功能正确;会看 profile 和理解 kernel,才能解释为什么慢以及怎么变快。
参考#
- Stanford CS336 Spring 2025 Course Materials: lecture_06.py