第 1 课 · 从手工雕琢到一行加速——每一层用多少"控制权"换多少"生产力"
GPU 上跑的任何计算最终都是CUDA kernel。但并非每个人都直接写 CUDA。 就像没人用汇编写 CRUD 应用——上层抽象让程序员专注于"做什么"而非"怎么做"。 下面是把三层叠在一起的全景图。从上往下是调用链;从下往上,每层拿控制权换生产力。
model = AutoLigerKernelForCausalLM.from_pretrained(...) —— 一行 patch 掉底层,用优化 kernel 替换默认实现。完全不需要懂 GPU。apply_liger_kernel_to_llama() 生效。6@triton.jit 写 kernel,block 级编程。编译器自动生成多线程调度、shared memory 分配、autotuning。torch.compile(Inductor) 的输出默认就是 Triton kernel。5__global__ + <<<>>> + thread/block/grid + shared/global memory。你控制一切——线程怎么排布、数据怎么搬、warp 怎么同步。这是第 2 课(线程层级)的核心内容。关键洞察:你在这条栈里怎么定位,取决于你的目标。 如果目标是优化一个已知的 LLM 算子(RoPE、RMSNorm),liger-kernel 已经帮你写好了——直接用。 如果目标是想理解 torch.compile 是怎么工作的,你需要知道它背后的 Inductor 默认输出 Triton 而非 CUDA C++。 如果目标是写出比库里带的内核更快的新算法——那就自己写 CUDA 或 Triton。
在第 2 课你会写一个让每个线程报"我是谁"的 kernel——代码很简单,但已经暴露了 CUDA 编程的所有核心复杂度:
| 你控制什么 | 如果搞错了 |
|---|---|
grid/block 大小(<<<B, T>>>) | occupancy 低下,SM 里一堆空闲 slot |
shared memory 分配与同步(__syncthreads()) | 数据竞争、死锁、bank conflict |
| 全局内存访问模式 | 非合并访问 → 带宽利用率从 ~90% 跌到 ~10%4 |
| 线程索引计算公式 | 越界读写、算错负责的数据 |
| warp 内分支路径 | warp divergence → 一半线程闲着 |
在 CUDA 的世界里,写一个"正确"的 kernel 只是及格线——写出"快"的 kernel 需要额外的 coalescing、shared tiling、register blocking、occupancy tuning……每一项都是独立的知识点。 Simon Boehm 的 matmul 优化日志一路从「一个线程算一个元素(naive)」优化到了 「1.3 TFLOPs 的缓存块 + 向量化」4——那是对 同一个数学运算 的十几轮迭代。
Triton(OpenAI 的开源项目,Philippe Tillet 等人 2019 年提出5)的核心创新: 你不需要为每一个线程写代码。
先看同一个 vector add 操作的 CUDA vs Triton 的代码量对比:
__global__ void vec_add(
float *x, float *y, float *o, int n
) {
int i = blockIdx.x * blockDim.x
+ threadIdx.x;
if (i < n) o[i] = x[i] + y[i];
}
// 你还要手动: cudaMalloc
// cudaMemcpy, launch config,
// block size 选多少……
@triton.jit
def vec_add(x_ptr, y_ptr, o_ptr,
n_elements,
BLOCK_SIZE: tl.constexpr):
pid = tl.program_id(0)
offs = pid*BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
mask = offs < n_elements
x = tl.load(x_ptr + offs, mask=mask)
y = tl.load(y_ptr + offs, mask=mask)
tl.store(o_ptr + offs, x + y, mask=mask)
| CUDA 里你手写的 | Triton 编译器自动搞定 |
|---|---|
| block/thread 维度选择(128? 256? 320?) | autotuning 自动搜索最优配置5 |
shared memory 分配 + __syncthreads() | 编译器自动插入——你只用标记哪些数据该放 shared |
| 全局内存 coalescing | 编译器分析访存模式自动对齐 |
边界条件 if (i < n) | mask=mask 自动处理越界 |
| 寄存器分配、指令调度 | LLVM 后端(Triton 编译链 → LLVM IR → PTX)5 |
核心哲学:你用较粗粒度的 program(概念上类似 CUDA 的 block)而非 thread 视角来写 kernel。
tl.program_id(0) 是"第几个程序"(类似 blockIdx),tl.arange(0, BLOCK_SIZE) 是"这个程序里的向量化工作单元"。
每个 program 内发生了什么——线程怎么切、数据怎么搬——编译器接管。7
tl.atomic_add,新版本加了 warp-specialization。但和非 reduction 的 element-wise 算子相比,reduction 在 Triton 里仍然需要更多手工。如果你已经理解了 CUDA(手写一切)和 Triton(编译器管线程),那 liger-kernel 就是一个 "请人帮你写了 Triton kernel" 的库。
LinkedIn 团队发现:LLM 训练中有几个高频算子——RMSNorm、RoPE、SwiGLU、CrossEntropy——占用了 大量显存和时间,但大家都在重复实现。于是他们把最优 Triton 实现开源了:
注意:liger-kernel 的 kernel 仍然是 Triton 写的。它的价值不在于"写了一种新语言",而在于 把 LLM 训练中最常见的优化模式封装好了。具体技法:
| 技法 | 做什么 | 为什么能省显存 |
|---|---|---|
| Kernel Fusion | 把 线性层 + CrossEntropy 合并成一个 kernel | 中间结果不写回 HBM,在寄存器 / shared memory 里直接消费6 |
| In-place 操作 | RMSNorm 直接覆盖输入 tensor | 省掉一份中间 buffer 的显存分配 |
| Chunking | 大 tensor 切成小块逐块算(像 FlashAttention 的 tiling) | 峰值显存只等于一个 chunk 大小,不是整个 tensor6 |
| Recomputation | 前向时不存中间激活,反向时重算 | 用计算换显存(和 FlashAttention 同款思路) |
和 FlashAttention 的相似点:FA 的核心优化——tiling + online softmax + recomputation——和 liger-kernel 的 chunking + in-place + fusion 是同一种哲学:在 kernel 里多做几步,少来回读 HBM。 FA 是 attention 专用优化;liger-kernel 把同样的思路搬到了 RMSNorm、RoPE、CrossEntropy 这些非 attention 算子上。6
# === 方式 1:一行 monkey-patch(最常用)===
from liger_kernel.transformers import apply_liger_kernel_to_llama
apply_liger_kernel_to_llama() # 替换 loss/RMSNorm/RoPE/SwiGLU
model = AutoModelForCausalLM.from_pretrained("meta-llama/Llama-3-8B")
# === 方式 2:AutoModel 自动 patch ===
from liger_kernel.transformers import AutoLigerKernelForCausalLM
model = AutoLigerKernelForCausalLM.from_pretrained("meta-llama/Llama-3-8B")
# === 方式 3:手拼 kernel(需要定制时)===
from liger_kernel.transformers import LigerFusedLinearCrossEntropyLoss
loss_fn = LigerFusedLinearCrossEntropyLoss()
loss = loss_fn(lm_head.weight, input, target)
loss.backward()
依赖极简:只需要 torch + triton(后面就是 CUDA),没有第三方库。8
这意味着你的训练环境只要装了 PyTorch + Triton,装 liger-kernel 就能用。
cp.async)有精确需求。现实是:绝大多数人的优化需求落在第 1 或第 2 条。第 3 条(必须手写 CUDA)通常是 NVIDIA 自己、FlashAttention 核心团队、或者硬件厂商在维持。 这也是为什么"学 CUDA 是为了理解,但用的主要是 Triton / liger-kernel"是一个务实的学习策略。
这篇 17 页的技术报告精确解释了 kernel fusion、chunking、in-place 三大武器的工作机制,以及它们在不同模型(LLaMA-8B)上带来 +20%/−60% 的 benchmark 数据。读它,把本课的表格对照着看。
Triton 文档的开篇介绍,用 tiled matmul 为例展示了"用 CUDA vs 用 Triton"的代码差异。读完你会直观感受到为什么 block 级编程是一种思维升级。