从零开始写Qwen3(四)实现RMSNorm算子
1. 概述
已经搭建了基本模型,可以推理,并且应用了KVCache,现在我们可以开始手写算子,先从最简单的RMSNorm开始:
1.1 什么是 RMS Norm?
每个特征向量除以方均根进行归一化,再乘以一个 gamma 进行尺度缩放
RMSNorm ( x ) i , j = x i , j ∑ j x i , j 2 n + ϵ ⋅ γ j \text{RMSNorm}(x)_{i,j} = \frac{x_{i,j}}{\sqrt{\frac{\sum_j x_{i,j}^2}{n}+ \epsilon}}\cdot \gamma_j RMSNorm(x)i,j=n∑jxi,j2+ϵxi,j⋅γj
2. 环境搭建
给torch写算子有两种方法:
- 直接在setup.py中写CUDAExtension/CPPExtension,然后python安装的时候会调用ninja(类似make的工具)编译cpp/cu文件,生成一个so文件
- 自己写CMakeLists.txt,自己配置python,pytorch等依赖,然后生成so文件,通过cmake的软链接到对应目录下
前者很方便,不用自己找torch、python、cuda等依赖,但缺点就是编译非常慢,它和python代码绑在一起,每次编译都需要刷新python的依赖
这里选择CMake项目,它编译快,灵活度高,整个项目配置一遍就不用管了
2.2 项目结构
qwen3_from_scratch/
├── kernels/
│ ├── rms_norm/
│ │ ├── rms_norm.cpp # CPU 实现 + Python 入口
│ │ └── rms_norm.cu # CUDA 实现
│ └── kernels.h
├── pybind11.cpp # 模块注册
├── CMakeLists.txt
└── cmake/
└── find_pytorch_vars.cmake
pybind11.cpp负责注册模块和所有的函数,一个目录一个算子,.cpp负责cpu实现,同时兼任算子入口,.cu负责GPU实现
2.3 CMake配置要点
- 使用 Python 获取 Torch/Python/CUDA 的 cmake 路径
- 设置 USE_CUDA 宏区分有/无 CUDA 环境
- 链接 torch_python 库解决符号未定义问题
- 编译后软链接到 Python 目录
详细代码可以参考代码仓中的CMakeLists.txt
3. 模块注册
在py11bind.cpp中写
#include "kernel.h"
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m)
{
m.doc() = "qwen3_from_scratch kernels";
m.def(
"rms_norm_forward",
&rms_norm_forward,
"RMSNorm forward computation (CPU/CUDA)",
py::arg("x"),
py::arg("gamma"),
py::arg("eps") = 1e-6f);
}
kernels.h中包含了<torch/extensions.h>,这里注册了一个叫做ops的模块,导出了一个叫做rms_norm_forward的函数,包含三个参数
4. CPU实现
cpu实现都比较简单,不做优化,说明原理,验证准确性就行
for (int i = 0; i < seqLen; i++)
{
const T *currentX = x + i * hiddenDimStride;
T *currentOutput = output + i * hiddenDim;
// 1. 计算当前序列位置的平方和
float sumSq = 0.0f;
for (int k = 0; k < hiddenDim; k++)
{
const float val = static_cast<float>(currentX[k]);
sumSq += val * val;
}
// 2. 计算均方根(RMS)
const float rms = sqrtf(sumSq / static_cast<float>(hiddenDim) + eps);
// 3. 归一化 + 缩放
for (int j = 0; j < hiddenDim; j++)
{
const float val = static_cast<float>(currentX[j]);
currentOutput[j] = static_cast<T>(val / rms * static_cast<float>(gamma[j]));
}
}
简单循环就行,外部把B、H、S合并成一个,也就是 reshape(-1, D),到这里就成为二维数据
5. CUDA实现
5.1 CUDA 并行层级回顾
- 线程束(warp),32个线程,同时执行相同指令,但每个线程有自己的寄存器,可以读取不同指令,也就是SIMD,单指令多数据
- 块(block),块包含多个线程束,最多1024个线程,块是调度基本单位,同时也是共享内存和线程同步的基本单位,最好不要跨块通信
- 网格(grid),块就是网格上的一个点,一次函数执行会把网格上所有点都执行完
5.2 计算步骤
- 求方均根(需要归约)
- 每个元素除以方均根(完全并行)
- 乘以 gamma(完全并行)
关键就在于步骤1的规约操作
5.3 规约优化策略
5.3.1 分层规约
首先根据线程块长度将向量分为多个部分,然后求和,每个线程一个,于是长度为N的向量变成了长度为整块线程数量的结果
for (uint32_t i = tid; i < hiddenDim; i += blockSize) {
sumSq += tempX[tempXj] * tempX[tempXj];
}
这样块中每个线程拥有一个变量sumSq,总共blockSize这么多个线程
然后下一步就要跨线程汇总,需要用到共享内存,把每个线程的值存进去,开始下一步汇总
跨线程汇总最典型的模式就是这样
for (int w = width / 2; w > 0 ; w /= 2) {
if (tid < w) {
smem[tid] += smem[tid + w];
}
__syncthreads();
}
这样每次把一半的元素累加,最后得到只剩一个数。
这种方法可以让参与计算的线程尽可能都在前几个warp中,而不是分散在各个warp中,可以降低线程分化
5.3.2 线程束打散规约
由于可能不止涉及一个线程束,必须等待所有线程束计算完毕,会产生大量等待消耗,cuda针对这种操作专门提供了线程束内的同步函数,使用这种函数就不需要线程束之间的同步,而且巧妙的是,一个线程块至多1024个线程,正好是32*32,所以进行一次线程束汇总后只剩下至多32个数据,只用放到一个线程束中执行一次线程束汇总就行,只需要一次同步
线程束内同步函数通常叫做 shfl_xxx_sync ,大致使用就是
shlf_xxx_sync(val, mask, offset)
val是单个变量
这就相当于是把整个线程束的val当成一个数组,第i个线程(i叫做 laneId )对应val[id]
if (((1 << i) && mask) && (xxx(i,offset)>=0 && xxx(i,offset)< 32) {
val[xxx(i, offset)]
}
常见的有
shfl_down_sync,就是xxx(i, offset)=i+offsetshlf_up_sync,类似shlf_xor_sync,就是xxx(i,offset)=i^offset
可以使用这个函数对32个线程进行无同步汇总
template <int width = WARP_SIZE, typename T> __device__ __forceinline__
T warp_reduce_sum(T x) {
#pragma unroll
for (int offset = width / 2; offset > 0; offset >>= 1) {
x += __shfl_xor_sync(0xffffffff, x, offset, width);
}
return x;
}
由于width已知,所以可以把循环展开成多条(5条)指令,把分支判断和跳转都去掉
为了方便展示,假设线程束只有8个线程,执行结果如下
初始: ['a0', 'a1', 'a2', 'a3', 'a4', 'a5', 'a6', 'a7']
step1: ['(a0+a4)', '(a1+a5)', '(a2+a6)', '(a3+a7)', '(a4+a0)', '(a5+a1)', '(a6+a2)', '(a7+a3)']
step2: ['(a0+a4+a2+a6)', '(a1+a5+a3+a7)', '(a2+a6+a0+a4)', '(a3+a7+a1+a5)', '(a4+a0+a6+a2)', '(a5+a1+a7+a3)', '(a6+a2+a4+a0)', '(a7+a3+a5+a1)']
step3: ['(a0+a4+a2+a6+a1+a5+a3+a7)', '(a1+a5+a3+a7+a0+a4+a2+a6)', '(a2+a6+a0+a4+a3+a7+a1+a5)', '(a3+a7+a1+a5+a2+a6+a0+a4)', '(a4+a0+a6+a2+a5+a1+a7+a3)', '(a5+a1+a7+a3+a4+a0+a6+a2)', '(a6+a2+a4+a0+a7+a3+a5+a1)', '(a7+a3+a5+a1+a6+a2+a4+a0)']
可以看到所有线程最后结果都变成一样,而down和up不是。一般这种操作只会输出线程0的结果
这样计算出来剩下不到32个有效值,用共享内存同步一次,再执行一次就行
__shared__ T s_sum[32];
const uint32_t warpId = tid / WARP_SIZE;
const uint32_t laneId = tid % WARP_SIZE;
if (laneId == 0) {
s_sum[warpId] = sumSq;
}
__syncthreads();
sumSq = 0.0f;
if (laneId < (blockSize / WARP_SIZE)) {
sumSq = s_sum[laneId];
}
sumSq = warp_reduce_sum(sumSq);
这里使用xor的好处就来了,xor版的warp_reduce_sum中所有线程拿到相同的值,就省去再获取一次值
5.4 完整流程
template <size_t blockSize, size_t hiddenDim>
__global__ void rms_norm_kernel_arr(const float* __restrict__ x,
float* __restrict__ output,
const float* __restrict__ gamma,
const int seqLen,
const int hiddenDimStride,
const float eps) {
uint32_t tid = threadIdx.x;
uint32_t blockId = blockIdx.x;
const T* x_ptr = x + blockId * hiddenDimStride; // 如果是连续,hiddenDimStride就是 hiddenDim
output += blockId * hiddenDim /* *1 output是刚申请的,stride肯定是1*/;
float sumSq = 0.0f;
#pragma unroll
for (uint32_t i = tid; i < hiddenDim; i += blockSize) {
float x = x_ptr[i];
sumSq += x * x;
}
sumSq = warp_reduce_sum(sumSq);
if constexpr (blockSize > WARP_SIZE) {
static_assert((blockSize <= 1024) && (blockSize % WARP_SIZE == 0), "blockSize must be a multiple of warpSize");
__shared__ float s_sum[32];
const uint32_t warpId = tid / WARP_SIZE;
const uint32_t laneId = tid % WARP_SIZE;
if (laneId == 0) {
s_sum[warpId] = sumSq;
}
__syncthreads();
sumSq = 0.0f;
if (laneId < (blockSize / WARP_SIZE)) {
sumSq = s_sum[laneId];
}
sumSq = warp_reduce_sum(sumSq);
}
const float mean = sumSq / hiddenDim;
const float scale = rsqrtf(mean + eps);
#pragma unroll
for (uint32_t i = tid; i < hiddenDim; i += blockSize) {
float x = x_ptr[i];
output[i] = static_cast<T>(x * scale * gamma[i]);
}
}
全程只需要一个核函数就可以解决
6. 性能测试与对比
6.1 测试环境
- GPU: RTX 3060
- 对比基准: PyTorch 2.10 nn.functional.rms_norm
- 测试数据: B×128×1024,多种数据类型
6.2 性能结果
| 数据类型 | 平均加速比 | 最佳加速比 |
|---|---|---|
| bfloat16 | 1.29x | 1.66x (Dim=128) |
| float16 | 1.20x | 1.50x (Dim=16384) |
| float32 | 1.29x | 1.67x (Dim=64) |
详情见 rms_norm_benchmark_report.md
6.3 torch融合算子对比
如果在torch2.8及以前的版本测试这个例子,会发现提升更加明显,甚至可以到几倍,因为torch2.8之前,nn.functional.rms_norm算子没有融合,是多个算子接连计算,性能大打折扣
可以看2.8版本
而2.10版本的
2.8版本是先调用pow(2),然后求均值,然后相加,再做除法,而2.10只有一个函数
6.4 CUDA加速总结
- 融合操作减少内存访问次数
- Warp shuffle 比共享内存归约更高效
- 减少 Python 层调用开销
更多推荐





所有评论(0)