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 = 4sqrt
token 0
RMS = sqrt(11) ≈ 3.317
token 1
RMS = sqrt(4) = 2normalize
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