从零开始写Qwen3目录

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=njxi,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配置要点

  1. 使用 Python 获取 Torch/Python/CUDA 的 cmake 路径
  2. 设置 USE_CUDA 宏区分有/无 CUDA 环境
  3. 链接 torch_python 库解决符号未定义问题
  4. 编译后软链接到 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 计算步骤

  1. 求方均根(需要归约)
  2. 每个元素除以方均根(完全并行)
  3. 乘以 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+offset
  • shlf_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加速总结

  1. 融合操作减少内存访问次数
  2. Warp shuffle 比共享内存归约更高效
  3. 减少 Python 层调用开销
Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐