LayerNorm(LN)

核心操作:

$$y = \frac{x-μ}{\sqrt{\frac{1}{d}\sum(x-μ)^2+ϵ}}γ+β$$

做了两件事:

  • 减均值(center):x−μ

  • 除标准差(scale):/σ

特点:

  • 输出是 零均值 + 单位方差

  • 有两个参数:γ(scale)+ β(bias)

  • 更“标准化”(distribution 更稳定)

RMSNorm

$$y = \frac{x}{\sqrt{\frac{1}{d}\sum{x^2}+ϵ}}γ$$

只做一件事:

  • 按 RMS 做缩放(scale)

  • 不减均值

LayerNorm = 去均值 + 归一化方差
RMSNorm  = 只归一化幅度(不关心均值)

现在的LLM结构里,LayerNorm层的权重名称上是LayerNorm,但是实际上,从数学公式上看大多使用的是RMSNorm

举例说明

如果一个大模型,他的hidden_size=5

我们构造一个简单例子:

  • hidden_size = 5

  • token 数 = 2

输入:

token0: [1, 2, 3, 4, 5]
token1: [2, 2, 2, 2, 2]

权重(γ):

weight = [2, 2, 2, 2, 2]

ε = 0(简化计算)

计算 mean(x²)

token 0
x² = [1,4,9,16,25]
sum = 55
mean = 55 / 5 = 11

token 1
x² = [4,4,4,4,4]
mean = 4

sqrt

token 0
RMS = sqrt(11) ≈ 3.317

token 1
RMS = sqrt(4) = 2

normalize

token 0
[1,2,3,4,5] / 3.317
≈ [0.30, 0.60, 0.90, 1.21, 1.51]

token 1
[2,2,2,2,2] / 2 = [1,1,1,1,1]

RMSNorm 最终结果

token0: [0.30, 0.60, 0.90, 1.21, 1.51]*[2, 2, 2, 2, 2]
=[0.6, 1.2, 1.8, 2.42, 3.02]
token1: [1, 1, 1, 1, 1]*[2, 2, 2, 2, 2]
=[2, 2, 2, 2, 2]

kernel实现

对一个shape为[5, 5120]的tensor进行rmsnorm归一化

Qwen3-32B模型的hidden_size为5120,这里我们模拟5个token的输入

线程组织方式:

如图1所示,线程由1个grid组成,1个grid里面包含5个block,每个block负责处理一个5120维的token向量,单线程单次global memory指令支持的访问宽度最大是16字节,本例中我们的数据类型是bfloat16,所以我们设计每个线程处理8个bfloat16数据,这样一个block就需要5120/8=640个线程,640就是block dim的大小,grid dim的大小为5,因为一grid里的5个block,我们这里grid和block都是一维的,可以想像成一个一维数组就行。

关键代码解析

int main()
{
    constexpr int num_tokens = 5;
    constexpr int hidden_size = 5120;
    constexpr int num_dims = 2;
    float epsilon = 1e-5;

    size_t size = num_tokens * hidden_size * sizeof(__nv_bfloat16);
    size_t w_size = hidden_size * sizeof(__nv_bfloat16);

    // host
    __nv_bfloat16* h_input = new __nv_bfloat16[num_tokens * hidden_size];
    __nv_bfloat16* h_weight = new __nv_bfloat16[hidden_size];
    __nv_bfloat16* h_out_gpu = new __nv_bfloat16[num_tokens * hidden_size];
    __nv_bfloat16* h_out_cpu = new __nv_bfloat16[num_tokens * hidden_size];

    // 初始化 input:0.2 ~ 1.0
    srand(time(nullptr));
    for (int i = 0; i < num_tokens * hidden_size; i++) {
        float r = static_cast<float>(rand()) / RAND_MAX;  // [0,1)
        float val = 0.2f + r * (1.0f - 0.2f);             // [0.2,1.0)
        h_input[i] = __float2bfloat16(val);
    }
    // 初始化 weight:-0.02 ~ 0.5
    for (int i = 0; i < hidden_size; i++) {
        float r = static_cast<float>(rand()) / RAND_MAX;   // [0,1)
        float val = -0.02f + r * (0.5f + 0.02f);           // [-0.02,0.5)
        h_weight[i] = __float2bfloat16(val);
    }

    // device
    __nv_bfloat16 *d_input, *d_weight, *d_out;
    cudaMalloc(&d_input, size);
    cudaMalloc(&d_weight, w_size);
    cudaMalloc(&d_out, size);

    cudaMemcpy(d_input, h_input, size, cudaMemcpyHostToDevice);
    cudaMemcpy(d_weight, h_weight, w_size, cudaMemcpyHostToDevice);

    // GPU
    rms_norm<__nv_bfloat16, num_tokens, hidden_size, num_dims>(d_out, d_input, d_weight, epsilon);

    cudaMemcpy(h_out_gpu, d_out, size, cudaMemcpyDeviceToHost);

    // CPU
    rms_norm_cpu(h_out_cpu, h_input, h_weight, epsilon,
                 num_tokens, hidden_size);

    // 打印每一行对比
    for (int t = 0; t < num_tokens; t++) {
        std::cout << "============================\n";
        print_row("GPU", h_out_gpu, t, hidden_size);
        print_row("CPU", h_out_cpu, t, hidden_size);
    }

    // 释放
    cudaFree(d_input);
    cudaFree(d_weight);
    cudaFree(d_out);
    delete[] h_input;
    delete[] h_weight;
    delete[] h_out_gpu;
    delete[] h_out_cpu;

    return 0;
}

上面是main函数,主要要负责创建输入输出,初始化数据,调用GPU和CPU计算函数进行rmsnorm的运算,最后对比GPU和CPU的结果。

template <typename DATA_TYPE, int NUM_TOKENS, int HIDDEN_SIZE, int NUM_DIM>
void rms_norm(DATA_TYPE* out,     // [num_tokens, hidden_size]
              DATA_TYPE* input,   // [num_tokens, hidden_size]
              DATA_TYPE* weight,  // [hidden_size]
              float epsilon) {
    std::cout << "-------------launch kernel-------------" <<std::endl;
    cudaStream_t stream;
    cudaStreamCreate(&stream);
    dim3 grid(NUM_TOKENS);
    const int VEC_SIZE =
          std::gcd(16 / sizeof(DATA_TYPE), HIDDEN_SIZE);
    std::cout << "VEC_SIZE = " << VEC_SIZE << std::endl;
    constexpr int BLOCK_SIZE = HIDDEN_SIZE / VEC_SIZE;
    dim3 block(BLOCK_SIZE);
    transformer::rms_norm_kernel<DATA_TYPE, VEC_SIZE, HIDDEN_SIZE, NUM_DIM, BLOCK_SIZE>
        <<<grid, block, 0, stream>>>(
            out,
            input,
            weight,
            epsilon
        );

    cudaStreamSynchronize(stream);

    cudaStreamDestroy(stream);
}

rms_norm这个函数负责进行kernel_launch,划分线程层次。

template <typename DATA_TYPE, int VEC_SIZE, int HIDDEN_SIZE, int NUM_DIMS, int BLOCK_SIZE>
__global__ void rms_norm_kernel(
    DATA_TYPE* __restrict__ out,           // [num_tokens, hidden_size]
    const DATA_TYPE* __restrict__ input,   // [num_tokens, hidden_size]
    const DATA_TYPE* __restrict__ weight,  // [hidden_size]
    const float epsilon) {
    __shared__ float s_variance;
    float variance = 0.0f;
    const DATA_TYPE* input_row;
    const int64_t input_stride_row = HIDDEN_SIZE;
    if constexpr (NUM_DIMS == 2) {
        // 2D for layernorm normal case [batch_size, hidden]
        input_row = input + blockIdx.x * input_stride_row;
    } else if constexpr (NUM_DIMS == 3) {
        // 3D
    } else if constexpr (NUM_DIMS == 4) {
        // 4D
    }

    auto vec_op = [&variance](const vec_n_t<DATA_TYPE, VEC_SIZE>& vec) {
        #pragma unroll
        for (int i = 0; i < VEC_SIZE; ++i) {
            float x = static_cast<float>(vec.val[i]);
            variance += x * x;
        }
    };
    auto scalar_op = [&variance](const DATA_TYPE& val) {
        float x = static_cast<float>(val);
        variance += x * x;
    };
}
  • 每一个线程需要知道自己要处理的数据的首地址在哪,线程在block内部,先算出当前block要处理的数据的首地址,input_stride_row表示每跨越一行,两个数据之间的间隔是多少,比input[0][0]和input[1][0]之间的stride就是一行数据的个数据,也就是HIDDEN_SIZE,同理input[0][1]和input[1][1]之间的stride也是HIDDEN_SIZE。对于第0个block来说,首地址是input,对于第1个block来说,首地址是input+1*5120。blockIdx.x表示block的下标。

input_row = input + blockIdx.x * input_stride_row;
  • vec_op一次处理8个BF16数据

  • scalar_op一次处理1个BF16数据

template <typename DATA_TYPE, size_t VEC_SIZE>
struct __align__(VEC_SIZE * sizeof(DATA_TYPE)) vec_n_t {
    DATA_TYPE val[VEC_SIZE];
};

template <int VEC_SIZE, typename DATA_TYPE, typename VEC_OP, typename SCALAR_OP>
__device__ inline void vectorize_read_with_alignment(const DATA_TYPE* input,
                                                    int hidden_size,
                                                    int tid, int stride,
                                                    VEC_OP&& vec_op,
                                                    SCALAR_OP&& scalar_op) {
    static_assert(VEC_SIZE > 0 && (VEC_SIZE & (VEC_SIZE - 1)) == 0,
                "VEC_SIZE must be a positive power-of-two");
    constexpr int WIDTH = VEC_SIZE * sizeof(DATA_TYPE);
    uintptr_t addr = reinterpret_cast<uintptr_t>(input);

    bool can_vec = ((addr & (WIDTH - 1)) == 0) && ((hidden_size & (VEC_SIZE - 1)) == 0);
    if (can_vec) {
        int num_vec = hidden_size / VEC_SIZE;

        using vin_t = vec_n_t<DATA_TYPE, VEC_SIZE>;
        auto* v_in = reinterpret_cast<const vin_t*>(input);

        for (int i = tid; i < num_vec; i += stride) {
            vin_t tmp = v_in[i];
            vec_op(tmp);
        }
        return;
    }

    int misalignment_offset = addr & (WIDTH - 1);
    int alignment_bytes = WIDTH - misalignment_offset;
    int prefix_elems = alignment_bytes & (WIDTH - 1);
    prefix_elems /= sizeof(DATA_TYPE);
    prefix_elems = min(prefix_elems, hidden_size);

    for (int i = tid; i < prefix_elems; i += stride) {
        scalar_op(input[i]);
    }

    input += prefix_elems;
    hidden_size -= prefix_elems;

    int num_vec = hidden_size / VEC_SIZE;
    using vin_t = vec_n_t<DATA_TYPE, VEC_SIZE>;
    auto* v_in = reinterpret_cast<const vin_t*>(input);

    for (int i = tid; i < num_vec; i += stride) {
        vec_op(v_in[i]);
    }

    int tail_start = num_vec * VEC_SIZE;
    for (int i = tid + tail_start; i < hidden_size; i += stride) {
        scalar_op(input[i]);
    }
}

vectorize_read_with_alignment 这是核心计算函数,主要做的就是向量化一次读取8个数据,然后完成计算,同时处理地址不对齐和数据个数据不对齐的情况

static_assert限制 VEC_SIZE 必须是 2 的幂

static_assert(VEC_SIZE > 0 && (VEC_SIZE & (VEC_SIZE - 1)) == 0,
                "VEC_SIZE must be a positive power-of-two");

VEC_SIZE 必须是 1,2,4,8,16...(2^n)

为什么 (x & (x - 1)) == 0 能判断

只要是 2 的幂,只有二进制的最高位是1,如8的二进制1000,"2的幂"和"2的幂 - 1"与运算的结果一定是0,反过也成立。

8  = 1000
7  = 0111
---------
&  = 0000

能完全通过向量化计算的条件是数据地址(转换成整数)必须是16字节的整数倍,数据个数必须是向量大小的整数倍(bf16时,一个向量大小为8个元素)

bool can_vec = ((addr & (WIDTH - 1)) == 0) && ((hidden_size & (VEC_SIZE - 1)) == 0);

一个vec_n_t就是一个包含8个BF16的数组

  • num_vec = 5120 / 8 = 640

  • input是16字节对齐的,所以可以直接reinterpret_cast重新解释成v_in,后面在处理v_in的时候,就会向量化的去处理

  • 对for循环的理解,要注意,这里每个线程只会执行一遍for循环的代码,比如当tid=0时,从v_in[0]的位置LOAD数据,赋值给vin_t类型的变量tmp,而vin_t就是vec_n_t<DATA_TYPE, VEC_SIZE>,这个数据结构表示8个BF16数据,v_in现在可以向量化访问,所以这时会从v_in[0]的位置一次LOAD 8个数据给tmp,然后调用vec_op完成对这8个数据的平方和运算,i += 640, i = 640,i >= num_vec,循环结束;当tid=1时,从v_in[1]的位置LOAD数据,v_in[0]代表前8个数据,v_in[1]表示下一批8个数据,这时会从v_in[1]的位置一次LOAD 8个数据给tmp,然后调用vec_op完成对这8个数据的平方和运算,i += 640,i = 641, i >= num_vec,循环结束;对于reinterpret_cast转换的理解,可以看图2

if (can_vec) {
        int num_vec = hidden_size / VEC_SIZE;

        using vin_t = vec_n_t<DATA_TYPE, VEC_SIZE>;
        auto* v_in = reinterpret_cast<const vin_t*>(input);

        for (int i = tid; i < num_vec; i += stride) {
            vin_t tmp = v_in[i];
            vec_op(tmp);
        }
        return;
    }

如图2所示,input每个均素表示一个2节字的BF16数据,而重新解释后的v_in每个元素可以看成一共16字节的8个BF16

下面的码处理数据首地址不是16B整数倍的情况

  • 假设数据首地址addr = 34, WIDTH=16B

  • misalignment_offset = 34& (16 - 1) = 33%16 = 2,表示数据首地址离16B整数倍差了2个字节

    • Tips:a % b == a & (b - 1),当b是2的幂时成立

  • alignment_bytes = 16 - 2 = 14,表式从34这个地址开始,到地址48为止,这14个字节要单独处理,不能向量化,从48开始才能向量化

  • prefix_elems = alignment_bytes % WIDTH = 14 % 16 = 14,prefix_elems表示要单独处理的字节数,但是上一步明明已经算出来alignment_bytes = 14,为什么不能用alignment_bytes呢?而要再算一次prefix_elems,为什么prefix_elems才是真正的要单独处理的数据个数?考虑一种情况,假如数据首地址addr = 32,则misalignment_offset=0,alignment_bytes=16-0=16,数据首地址是对齐的,结果算出来alignment_bytes=16,alignment_bytes表示要单独处理,不能向量化的字节数,而这个16个数据事实上是可以向量化的,所以我们再计算一个prefix_elems = alignment_bytes & (WIDTH - 1) = 16 % 16 = 0,这下表示没有需要单独处理的元素了。

  • prefix_elems /= sizeof(DATA_TYPE),表示把字节数转换成对应的数据个数

  • prefix_elems = min(prefix_elems, hidden_size),处理真实数据不够的情况,在这个例子里,首地址addr = 34,表示有14个字节的数据要单独处理,14个字节对7个BF16数据,假如hidden_size=2,那根本没有7个数据需要处理,所以要和真实数据求一个min

int misalignment_offset = addr & (WIDTH - 1);
int alignment_bytes = WIDTH - misalignment_offset;
int prefix_elems = alignment_bytes & (WIDTH - 1);
prefix_elems /= sizeof(DATA_TYPE);
prefix_elems = min(prefix_elems, hidden_size);

下面的代码就比较好理解了,先单独处理prefix_elems个元数,然后input地址往后偏移到对齐的地址,hidden_size表示实际还剩下的数据个数,接着通过向量化完成剩下的计算。如果说剩下的hidden_size个数据个数据不是VEC_SIZE的整数倍,那么就是有尾部数据,也需要单独处理,这个也比较好理解就不赘述了。

    for (int i = tid; i < prefix_elems; i += stride) {
        scalar_op(input[i]);
    }

    input += prefix_elems;
    hidden_size -= prefix_elems;

    int num_vec = hidden_size / VEC_SIZE;
    using vin_t = vec_n_t<DATA_TYPE, VEC_SIZE>;
    auto* v_in = reinterpret_cast<const vin_t*>(input);

    for (int i = tid; i < num_vec; i += stride) {
        vec_op(v_in[i]);
    }
    
    int tail_start = num_vec * VEC_SIZE;
    for (int i = tid + tail_start; i < hidden_size; i += stride) {
        scalar_op(input[i]);
    }

上面的代码,每个线程完成了8个数据的平方和的计算,接着通过reduce完成一个block内部5120个数的平方和。再求标准差的倒数,标准差的倒数写入到共享内存s_variance,在__syncthreads之后,一个block内部的所有的线程都能看到这个值。最后每个线程把自己负责的那一部分数据乘上s_variance再乘上weight,把结果写回到global memory。

using BlockReduce = cub::BlockReduce<float, BLOCK_SIZE>;
    __shared__ typename BlockReduce::TempStorage reduceStore;
    variance = BlockReduce(reduceStore).Sum(variance);

    if (threadIdx.x == 0) {
        s_variance = rsqrtf(variance / HIDDEN_SIZE + epsilon);
    }
    __syncthreads();

    DATA_TYPE* out_row = out + blockIdx.x * HIDDEN_SIZE;
    auto* v_in = reinterpret_cast<const vec_n_t<DATA_TYPE, VEC_SIZE>*>(input_row);
    auto* v_w = reinterpret_cast<const vec_n_t<DATA_TYPE, VEC_SIZE>*>(weight);
    auto* v_out = reinterpret_cast<vec_n_t<DATA_TYPE, VEC_SIZE>*>(out_row);
    for (int i = threadIdx.x; i < HIDDEN_SIZE / VEC_SIZE; i += blockDim.x) {
        vec_n_t<DATA_TYPE, VEC_SIZE> dst;
        vec_n_t<DATA_TYPE, VEC_SIZE> src1 = v_in[i];
        vec_n_t<DATA_TYPE, VEC_SIZE> src2 = v_w[i];
        #pragma unroll
        for (int j = 0; j < VEC_SIZE; j++) {
            float x = static_cast<float>(src1.val[j]);
            dst.val[j] = ((DATA_TYPE)(x * s_variance)) * src2.val[j];
        }
        v_out[i] = dst;
    }

以上就是vLLM中rmsnorm的kernel实现,需要注意的是rmsnorm有两种实现,一种是不带残差的,就是本文中讲的这一种,另一种是带残差的,理解了不带的残差的rmsnorm,就能很轻松理解带残差的kernel。

完整源码

#include <cub/cub.cuh>
#include <cuda_bf16.h>
#include <numeric>
#include <cmath>
#include <iostream>
#include <cstdlib>
#include <ctime>

namespace transformer {
template <typename DATA_TYPE, size_t VEC_SIZE>
struct __align__(VEC_SIZE * sizeof(DATA_TYPE)) vec_n_t {
    DATA_TYPE val[VEC_SIZE];
};

template <int VEC_SIZE, typename DATA_TYPE, typename VEC_OP, typename SCALAR_OP>
__device__ inline void vectorize_read_with_alignment(const DATA_TYPE* input,
                                                    int hidden_size,
                                                    int tid, int stride,
                                                    VEC_OP&& vec_op,
                                                    SCALAR_OP&& scalar_op) {
    static_assert(VEC_SIZE > 0 && (VEC_SIZE & (VEC_SIZE - 1)) == 0,
                "VEC_SIZE must be a positive power-of-two");
    constexpr int WIDTH = VEC_SIZE * sizeof(DATA_TYPE);
    uintptr_t addr = reinterpret_cast<uintptr_t>(input);

    bool can_vec = ((addr & (WIDTH - 1)) == 0) && ((hidden_size & (VEC_SIZE - 1)) == 0);
    if (can_vec) {
        int num_vec = hidden_size / VEC_SIZE;

        using vin_t = vec_n_t<DATA_TYPE, VEC_SIZE>;
        auto* v_in = reinterpret_cast<const vin_t*>(input);

        for (int i = tid; i < num_vec; i += stride) {
            vin_t tmp = v_in[i];
            vec_op(tmp);
        }
        return;
    }

    int misalignment_offset = addr & (WIDTH - 1);
    int alignment_bytes = WIDTH - misalignment_offset;
    int prefix_elems = alignment_bytes & (WIDTH - 1);
    prefix_elems /= sizeof(DATA_TYPE);
    prefix_elems = min(prefix_elems, hidden_size);

    for (int i = tid; i < prefix_elems; i += stride) {
        scalar_op(input[i]);
    }

    input += prefix_elems;
    hidden_size -= prefix_elems;

    int num_vec = hidden_size / VEC_SIZE;
    using vin_t = vec_n_t<DATA_TYPE, VEC_SIZE>;
    auto* v_in = reinterpret_cast<const vin_t*>(input);

    for (int i = tid; i < num_vec; i += stride) {
        vec_op(v_in[i]);
    }

    int tail_start = num_vec * VEC_SIZE;
    for (int i = tid + tail_start; i < hidden_size; i += stride) {
        scalar_op(input[i]);
    }
}

template <typename DATA_TYPE, int VEC_SIZE, int HIDDEN_SIZE, int NUM_DIMS, int BLOCK_SIZE>
__global__ void rms_norm_kernel(
    DATA_TYPE* __restrict__ out,           // [num_tokens, hidden_size]
    const DATA_TYPE* __restrict__ input,   // [num_tokens, hidden_size]
    const DATA_TYPE* __restrict__ weight,  // [hidden_size]
    const float epsilon) {
    __shared__ float s_variance;
    float variance = 0.0f;
    const DATA_TYPE* input_row;
    const int64_t input_stride_row = HIDDEN_SIZE;
    if constexpr (NUM_DIMS == 2) {
        // 2D for layernorm normal case [batch_size, hidden]
        input_row = input + blockIdx.x * input_stride_row;
    } else if constexpr (NUM_DIMS == 3) {
        // 3D
    } else if constexpr (NUM_DIMS == 4) {
        // 4D
    }

    auto vec_op = [&variance](const vec_n_t<DATA_TYPE, VEC_SIZE>& vec) {
        #pragma unroll
        for (int i = 0; i < VEC_SIZE; ++i) {
            float x = static_cast<float>(vec.val[i]);
            variance += x * x;
        }
    };
    auto scalar_op = [&variance](const DATA_TYPE& val) {
        float x = static_cast<float>(val);
        variance += x * x;
    };
    vectorize_read_with_alignment<VEC_SIZE>(
        input_row, HIDDEN_SIZE, threadIdx.x, blockDim.x, vec_op, scalar_op);

    using BlockReduce = cub::BlockReduce<float, BLOCK_SIZE>;
    __shared__ typename BlockReduce::TempStorage reduceStore;
    variance = BlockReduce(reduceStore).Sum(variance);

    if (threadIdx.x == 0) {
        s_variance = rsqrtf(variance / HIDDEN_SIZE + epsilon);
    }
    __syncthreads();

    DATA_TYPE* out_row = out + blockIdx.x * HIDDEN_SIZE;
    auto* v_in = reinterpret_cast<const vec_n_t<DATA_TYPE, VEC_SIZE>*>(input_row);
    auto* v_w = reinterpret_cast<const vec_n_t<DATA_TYPE, VEC_SIZE>*>(weight);
    auto* v_out = reinterpret_cast<vec_n_t<DATA_TYPE, VEC_SIZE>*>(out_row);
    for (int i = threadIdx.x; i < HIDDEN_SIZE / VEC_SIZE; i += blockDim.x) {
        vec_n_t<DATA_TYPE, VEC_SIZE> dst;
        vec_n_t<DATA_TYPE, VEC_SIZE> src1 = v_in[i];
        vec_n_t<DATA_TYPE, VEC_SIZE> src2 = v_w[i];
        #pragma unroll
        for (int j = 0; j < VEC_SIZE; j++) {
            float x = static_cast<float>(src1.val[j]);
            dst.val[j] = ((DATA_TYPE)(x * s_variance)) * src2.val[j];
        }
        v_out[i] = dst;
    }
}

}

template <typename DATA_TYPE, int NUM_TOKENS, int HIDDEN_SIZE, int NUM_DIM>
void rms_norm(DATA_TYPE* out,     // [num_tokens, hidden_size]
              DATA_TYPE* input,   // [num_tokens, hidden_size]
              DATA_TYPE* weight,  // [hidden_size]
              float epsilon) {
    std::cout << "-------------launch kernel-------------" <<std::endl;
    cudaStream_t stream;
    cudaStreamCreate(&stream);
    dim3 grid(NUM_TOKENS);
    const int VEC_SIZE =
          std::gcd(16 / sizeof(DATA_TYPE), HIDDEN_SIZE);
    std::cout << "VEC_SIZE = " << VEC_SIZE << std::endl;
    constexpr int BLOCK_SIZE = HIDDEN_SIZE / VEC_SIZE;
    dim3 block(BLOCK_SIZE);
    transformer::rms_norm_kernel<DATA_TYPE, VEC_SIZE, HIDDEN_SIZE, NUM_DIM, BLOCK_SIZE>
        <<<grid, block, 0, stream>>>(
            out,
            input,
            weight,
            epsilon
        );

    cudaStreamSynchronize(stream);

    cudaStreamDestroy(stream);
}

// CPU RMSNorm
void rms_norm_cpu(__nv_bfloat16* out,
                  __nv_bfloat16* input,
                  __nv_bfloat16* weight,
                  float epsilon,
                  const int num_tokens,
                  const int hidden_size) {
    for (int t = 0; t < num_tokens; t++) {
        float variance = 0.0f;

        // 1. 计算平方和
        for (int i = 0; i < hidden_size; i++) {
            float x = __bfloat162float(input[t * hidden_size + i]);
            variance += x * x;
        }

        // 2. 计算 rsqrt
        float scale = rsqrtf(variance / hidden_size + epsilon);

        // 3. 归一化 + weight
        for (int i = 0; i < hidden_size; i++) {
            float x = __bfloat162float(input[t * hidden_size + i]);
            float w = __bfloat162float(weight[i]);
            float y = x * scale * w;
            out[t * hidden_size + i] = __float2bfloat16(y);
        }
    }
}

// 打印一行的前10个元素和后10个元素
void print_row(const char* name,
               __nv_bfloat16* data,
               int row,
               int hidden_size) {
    std::cout << name << " row " << row << ":\n";

    std::cout << "  first 10: ";
    for (int i = 0; i < 10; i++) {
        float v = __bfloat162float(data[row * hidden_size + i]);
        std::cout << v << " ";
    }

    std::cout << "\n  last 10: ";
    for (int i = hidden_size - 10; i < hidden_size; i++) {
        float v = __bfloat162float(data[row * hidden_size + i]);
        std::cout << v << " ";
    }
    std::cout << "\n";
}

int main()
{
    std::cout << "-------------main-------------" <<std::endl;
    // out shape[5, 5120]
    // input shape[5, 5120]
    // weight shpae[5120]
    // epsilon 0.00001
    constexpr int num_tokens = 5;
    constexpr int hidden_size = 5120;
    //constexpr int hidden_size = 34;
    constexpr int num_dims = 2;
    float epsilon = 1e-5;

    size_t size = num_tokens * hidden_size * sizeof(__nv_bfloat16);
    size_t w_size = hidden_size * sizeof(__nv_bfloat16);

    // host
    __nv_bfloat16* h_input = new __nv_bfloat16[num_tokens * hidden_size];
    __nv_bfloat16* h_weight = new __nv_bfloat16[hidden_size];
    __nv_bfloat16* h_out_gpu = new __nv_bfloat16[num_tokens * hidden_size];
    __nv_bfloat16* h_out_cpu = new __nv_bfloat16[num_tokens * hidden_size];

    // 初始化 input:0.2 ~ 1.0
    srand(time(nullptr));
    for (int i = 0; i < num_tokens * hidden_size; i++) {
        float r = static_cast<float>(rand()) / RAND_MAX;  // [0,1)
        float val = 0.2f + r * (1.0f - 0.2f);             // [0.2,1.0)
        h_input[i] = __float2bfloat16(val);
    }
    // 初始化 weight:-0.02 ~ 0.5
    for (int i = 0; i < hidden_size; i++) {
        float r = static_cast<float>(rand()) / RAND_MAX;   // [0,1)
        float val = -0.02f + r * (0.5f + 0.02f);           // [-0.02,0.5)
        h_weight[i] = __float2bfloat16(val);
    }

    // device
    __nv_bfloat16 *d_input, *d_weight, *d_out;
    cudaMalloc(&d_input, size);
    cudaMalloc(&d_weight, w_size);
    cudaMalloc(&d_out, size);

    cudaMemcpy(d_input, h_input, size, cudaMemcpyHostToDevice);
    cudaMemcpy(d_weight, h_weight, w_size, cudaMemcpyHostToDevice);

    // GPU
    rms_norm<__nv_bfloat16, num_tokens, hidden_size, num_dims>(d_out, d_input, d_weight, epsilon);

    cudaMemcpy(h_out_gpu, d_out, size, cudaMemcpyDeviceToHost);

    // CPU
    rms_norm_cpu(h_out_cpu, h_input, h_weight, epsilon,
                 num_tokens, hidden_size);

    // 打印每一行对比
    for (int t = 0; t < num_tokens; t++) {
        std::cout << "============================\n";
        print_row("GPU", h_out_gpu, t, hidden_size);
        print_row("CPU", h_out_cpu, t, hidden_size);
    }

    // 释放
    cudaFree(d_input);
    cudaFree(d_weight);
    cudaFree(d_out);
    delete[] h_input;
    delete[] h_weight;
    delete[] h_out_gpu;
    delete[] h_out_cpu;

    return 0;
}
nvcc rmsnorm.cu -o rmsnorm   -std=c++17   -O3   -arch=sm_90  -I/usr/local/cuda/include -L/usr/local/cuda/lib64 -lcuda

./rmsnorm