从零开始写Qwen3(四-其二)使用Triton实现RMSNorm算子
概述
前文用CUDA实现了RMSNorm这个算子,这个算子还算简单的,但还需要自己手动管理各种内容,还需要额外编译,使用pytorch调用时还需要先加载so,然后用动态链接库调用函数
本文使用 triton 这个框架来用 python 来实现同等功能,在大大简化开发的情况下性能额不输 CUDA 的版本
Triton RMSNorm的基本实现
参照前文,可以写出 Triton 版本
@triton.jit
def rms_norm_forward_kernel(
x,
gamma,
output,
stride,
eps: float,
D: tl.constexpr,
BLOCK_SIZE: tl.constexpr,
):
pid = tl.program_id(0)
tid = tl.arange(0, BLOCK_SIZE)
x += pid * stride
output += pid * stride
sum_sq = tl.zeros((BLOCK_SIZE,), tl.float32)
for block_offset in tl.range(0, D, BLOCK_SIZE):
cols = block_offset + tid
mask = cols < D
items = tl.load(x + cols, mask, 0.0).to(tl.float32)
sum_sq += items * items
var = tl.sum(sum_sq)
scale = tl.rsqrt(var / D + eps)
for block_offset in tl.range(0, D, BLOCK_SIZE):
cols = block_offset + tid
mask = cols < D
items = tl.load(x + cols, mask, 0.0)
gamma_items = tl.load(gamma + cols, mask, 0.0)
tl.store(output + cols, items * scale * gamma_items, mask)
triton直接把线程块内并行直接实现了,不需要开发者手动实现,比如这个 tl.sum(sum_sq)
另外可以看到这些 x 都是读取了两次的,第一次读取用来计算标准差,第二次用来作为分子,如果 D 不是很大,可以把整行向量全部放到寄存器中,只用加载一次,有一定的提速
@triton.jit
def rms_norm_forward_kernel_one_step(
x,
gamma,
output,
stride,
eps: float,
D: tl.constexpr,
BLOCK_SIZE: tl.constexpr,
):
tl.static_assert(BLOCK_SIZE >= D)
pid = tl.program_id(0)
tid = tl.arange(0, BLOCK_SIZE)
x += pid * stride
output += pid * stride
cols = tid
mask = cols < D
x_items = tl.load(x + cols, mask, 0.0)
sum_sq = x_items * x_items
var = tl.sum(sum_sq)
scale = tl.rsqrt(var / D + eps)
gamma_items = tl.load(gamma + cols, mask, 0.0)
tl.store(output + cols, x_items * scale * gamma_items, mask)
这样要求 BLOCK_SIZE>=D
性能对比
使用triton.testing.perf_report 进行测试,对比 torch 版本(nn.functional.rms_norm),前一章的 CUDA 版,用的实际上是上面第二种寄存器缓存数据版本,BLOCK_SIZE和D成正比,到1024截止,以及本章的 triton,数据分为 FP32,FP16 和 BF16,统一使用 2x256xD 的数据,D从 128 到 8196
triton使用第二种,直接使用 BLOCK_SIZE=triton.next_power_of_2(D),而不写 num_warps,默认使用128线程
FP32
FP16
BF16
CUDA和triton没有采用相同的结构,正好可以验证BLOCK_SIZE和线程数量的影响:CUDA的线程数量就等于BLOCK_SIZE,没有内部循环,在2176这里性能出现大幅下跌,因为它只有2176个数据,却要用4096个进行计算,有将近50%都有浪费。而triton使用128线程,虽然外部循环还是2048循环,循环两次,但内部并行是128线程,只需要多执行一次 ,没有出现大量浪费
torch的版本显然是固定BLOCK_SIZE,随着D的增长,性能平稳甚至有些下降
对比triton
- triton的BLOCK_SIZE完全可以和线程数量不同,它会自动实现内部循环,而CUDA需要自己手动实现
- triton的全局内存读写都是优化过的,使用前置谓词,避免了边界情况出现的if-else线程束分化,CUDA也需要写,不过可能要用到汇编,比如CUTLASS库中全局内存读写用的就是汇编
template <typename AccessType>
struct global_load<AccessType,
4,
CacheOperation::Always
> {
CUTLASS_DEVICE
global_load(AccessType &D, void const *ptr, bool pred_guard) {
unsigned &data = reinterpret_cast<unsigned &>(D);
asm volatile(
"{\n"
" .reg .pred p;\n"
" setp.ne.b32 p, %2, 0;\n"
" mov.b32 %0, %3;\n"
#if CUTLASS_ENABLE_L2_PREFETCH
" @p ld.global.L2::128B.u32 %0, [%1];\n"
#else
" @p ld.global.u32 %0, [%1];\n"
#endif
"}\n"
: "=r"(data)
: "l"(ptr), "r"((int)pred_guard), "r"(data));
}
};
重点在这里@p ld.global.u32 %0, [%1];,这个意思是如果p成立,则执行ld.global.u32,也就是加载全局内存到%0这个寄存器,也就是data这个变量上,如果不成立则什么也不做
和if-else的不同,这里没有任何跳转指令,满足条件和不满足条件的执行的完全是相同的指令,没有任何线程束分化
- triton内部处理了不同的类别,这一点CUDA也需要自己手写,比如FP32,FP16,BF16用triton一套代码就能解决,CUDA则需要使用模板函数
- triton是运行时编译,这意味着它可以把一些运行时才能得到的东西当成常量,比如 BLOCK_SIZE,它和 D 相关,在CUDA中只能用if-else硬编码一些D出来
总之triton能实现的CUDA也都能实现,但triton实现大大简化,而且性能完全可以和CUDA媲美
踩坑点

一开始跑出来的结果Triton差很多,定位了很久,甚至都到汇编里看了,看不出任何问题,用ncu单个算子测量性能并没有比torch差,用nsys测量循环多次triton也是耗时最少的
最后突然意识到:可能不是GPU算子的问题,在调用算子之前的部分,然后就找到了这个
output = torch.zeros_like(x)
这个不仅是开辟显存空间,还全部初始化为0,就是这一步拖累了triton。实际上根本不需要全部写0,因为这个数据本来就是要写入的,所以只需要开辟空间就行,使用
output = torch.empty_like(x)
改成这样之后triton的性能立刻超过torch到达CUDA水平
另外一个踩坑点就是
sum_sq = tl.zeros((BLOCK_SIZE,), tl.float32)
for block_offset in tl.range(0, D, BLOCK_SIZE):
cols = block_offset + tid
mask = cols < D
items = tl.load(x + cols, mask, 0.0).to(tl.float32)
sum_sq += items * items
var = tl.sum(sum_sq)
这里,如果 sum_sq 使用 float,每个内部循环先执行一次tl.sum,性能也会大打折扣,因为 reduce 操作是需要显式等待同一个线程块所有线程束执行才能进行的,如果先求和,会在每次循环都执行reduce操作,浪费大量时间在等待同步上
更多推荐




所有评论(0)