CUDA算子优化(2):RMSNorm

1. 简介

RMSNorm和Batch Normalization(BN)、Layer Normalization(LN)等方法一样,都属于一种归一化方法,是提升训练稳定性、加速收敛的重要技巧之一。这里,我们不讨论这些方法的优劣性,我们只关注RMSNorm是如何计算的,以及CUDA程序如何实现和优化。

RMSNorm仅使用均方根归一化,即计算数据的平方的平均值,再开平方根,用来衡量数据的整体大小。核心公式为输入数据除以他们的均方根,然后乘以一个可学习的缩放参数 $\gamma$ 。具体公式可表示为:

$$
\bar{x}_i = \frac{x_i}{\mathrm{RMS}(x)} \cdot \gamma_i
$$

其中,$\mathrm{RMS}(x) = \sqrt{\frac{1}{C} \sum_{i=1}^{C} x_i^2+ \epsilon}$

2. 算子CPU实现

假设我们要处理的序列维度是B=8,T=1024,C=4096,其中B是batch_size,T是序列的长度,C是通道数;要求均方根,是在C这个维度进行求解。假设h_data为输入,维度为B*T*C;h_weights为缩放因子,维度为C;h_out为输出,维度为B*T*C。我们把计算过程分为三个阶段:

  1. 计算每一行的平方和
  2. 计算RMS(x)
  3. 乘上缩放因子,计算最终结果
void rmsnorm_cpu(float *h_out, float *h_data, float *h_weights, int B, int T, int C) {
for (int i = 0; i < B * T; ++i) {
float *inp = h_data + i * C;
float *out = h_out + i * C;

float local_sum = 0.0f;
for (int j = 0; j < C; ++j) {
local_sum += inp[j] * inp[j];
}
local_sum = 1 / sqrt(local_sum / C + 1e-6);
for (int j = 0; j < C; ++j) {
out[j] = inp[j] * local_sum * h_weights[j];
}
}
}

3. CUDA版本V0一个线程处理一行

假设一个线程处理一行,那么总共需要B*T个线程,以block_size=64为例,grid_size为(B*T + block_size - 1) / block_size。代码如下:

#include <stdio.h>
#include <stdlib.h>

#define BLOCK_SIZE 512

__global__ void kernel(float *d_out, float *d_data, float *weights, int B, int T, int C) {
int tid = threadIdx.x + blockDim.x * blockIdx.x;
int N = B * T;

if (tid < N) {
float *inp = d_data + tid * C;
float *out = d_out + tid * C;

float local_sum = 0.0f;
for (int i = 0; i < C; ++i) {
local_sum += inp[i] * inp[i];
}
local_sum = sqrt(local_sum / C + 1e-9);

for (int i = 0; i < C; ++i) {
out[i] = inp[i] / local_sum * weights[i];
}
}
}

void rmsnorm_cpu(float *h_out, float *h_data, float *weights, int B, int T, int C) {
for (int row = 0; row < B * T; ++row) {
float *inp = h_data + row * C;
float *out = h_out + row * C;
float local_sum = 0.0f;
for (int i = 0; i < C; ++i) {
local_sum += inp[i] * inp[i];
}
local_sum = 1 / sqrt(local_sum / C + 1e-6);
for (int i = 0; i < C; ++i) {
out[i] = local_sum * inp[i] * weights[i];
}
}
}

int check_results(const float *actual, const float *expected,
int n, float tol) {
int errors = 0;
float max_diff = 0.0f;
for (int i = 0; i < n; i++) {
float diff = fabsf(actual[i] - expected[i]);
if (diff > max_diff)
max_diff = diff;
if (diff > tol)
errors++;
}
printf(" max_diff = %.2e\n", max_diff);
printf(" GPU vs CPU : %s\n", errors == 0 ? "PASS" : "FAIL");
return errors == 0;
}

void init(float* h_data, int N, int low, int high) {
for (int i = 0; i < N; ++i) {
float r = (float) rand() / RAND_MAX;
h_data[i] = low + (high - low) * r;
}
}

int main() {
int B = 64, T = 1024, C = 2048;
int N = B * T * C;
size_t bytes = N * sizeof(float);
float *h_data = (float *)malloc(bytes);
float *h_out = (float *)malloc(bytes);
float *out = (float *)malloc(bytes);
float *weights = (float *)malloc(C * sizeof(float));

init(h_data, N, -1.0, 1.0);
init(weights, C, 0.0, 1.0);

float *d_data, *d_out, *d_weights;

cudaMalloc(&d_data, bytes);
cudaMalloc(&d_out, bytes);
cudaMalloc(&d_weights, C * sizeof(float));
cudaMemcpy(d_data, h_data, bytes, cudaMemcpyHostToDevice);
cudaMemcpy(d_weights, weights, C * sizeof(float), cudaMemcpyHostToDevice);

dim3 grid_size = (B*T + BLOCK_SIZE - 1) / BLOCK_SIZE;
kernel<<<grid_size, BLOCK_SIZE>>>(d_out, d_data, d_weights, B, T, C);
cudaMemcpy(out, d_out, bytes, cudaMemcpyDeviceToHost);

rmsnorm_cpu(h_out, h_data, weights, B, T, C);

check_results(h_out, out, B * T * C, 1e-5);


return 0;
}

4. CUDA版本V1一个Warp处理一行

很明显,可以看出,使用一个线程处理一行访存是不合并的,所以我们可以使用一个Warp处理一行,减少访存不合并的现象;总共需要B*T*warp_size个线程。我们仍然将其分为三个阶段。

  1. 计算每一行的平方和。相比于使用共享内存,我们可以直接使用__shfl_down_sync函数,这个函数直接在寄存器层面完成Warp内线程间数据交换,完全规避了共享内存的硬件延迟、同步开销和Bank Conflict的风险,速度更快。这里介绍一下__shfl_down_sync和__shfl_sync函数

__shfl_down_sync(mask, v, d, w),标号为t的参与线程返回标号为t+d的线程中变量v的值,标号满足t+d>=w的线程返回原来的v。例如,当w=8,d=2时,该函数将第2-7号线程中变量v的值传递给第0-5号线程,而第6-7号线程返回它们原来的v。形象地说,这是一种将数据向下平移的操作。

__shfl_sync(mask, v, srcLane, w),参与线程返回标号为 srcLane 的线程中变量 v 的值。这是一种广播式数据交换:将一个线程中的数据广播到所有(包括自己)线程。

Warp间数据规约的代码如下:

unsigned mask = 0xFFFFFFFF;
for (int offset = BLOCK_SIZE / 2; offset >= 1; offset /= 2) {
local_sum += __shfl_down_sync(mask, local_sum, offset);
}

local_sum = __shfl_sync(mask, local_sum, 0);
  1. 计算RMS(x)
  2. 乘上缩放因子,计算最终结果
    全部代码如下:
#include <stdio.h>
#include <stdlib.h>


#define BLOCK_SIZE 32

__global__ void kernel(float *d_out, float *d_data, float *d_weights, int N, int C) {
int index = threadIdx.x + blockDim.x * blockIdx.x;
int warp_id = index / blockDim.x;
int lane_id = index % blockDim.x;

if (warp_id > N) {
return;
}
float *inp = d_data + warp_id * C;
float *out = d_out + warp_id * C;

float local_sum = 0.0f;
for (int stride = lane_id; stride < C; stride += BLOCK_SIZE) {
local_sum += inp[stride] * inp[stride];
}

unsigned mask = 0xFFFFFFFF;
for (int offset = BLOCK_SIZE / 2; offset >= 1; offset /= 2) {
local_sum += __shfl_down_sync(mask, local_sum, offset);
}

local_sum = __shfl_sync(mask, local_sum, 0);
local_sum = 1 / sqrt(local_sum / C + 1e-6);
for (int i = 0; i < C; ++i) {
out[i] = inp[i] * local_sum * d_weights[i];
}

}

void init(float *data, int N, int low, int high) {
for (int i = 0; i < N; ++i) {
float r = (float)rand() / RAND_MAX;
data[i] = low + (high - low) * r;
}
}

void rmsnorm_cpu(float *h_out, float *h_data, float *weights, int N, int C) {
for (int i = 0; i < N; ++i) {
float *inp = h_data + i * C;
float *out = h_out + i * C;

float local_sum = 0.0f;
for (int j = 0; j < C; ++j) {
local_sum += inp[j] * inp[j];
}
local_sum = 1 / sqrt(local_sum / C + 1e-6);
for (int j = 0; j < C; ++j) {
out[j] = inp[j] * local_sum * weights[j];
}
}
}

int check_results(const float *actual, const float *expected,
int n, float tol) {
int errors = 0;
float max_diff = 0.0f;
for (int i = 0; i < n; i++) {
float diff = fabsf(actual[i] - expected[i]);
if (diff > max_diff)
max_diff = diff;
if (diff > tol)
errors++;
}
printf(" max_diff = %.2e\n", max_diff);
printf(" GPU vs CPU : %s\n", errors == 0 ? "PASS" : "FAIL");
return errors == 0;
}

int main() {
int B = 64, T = 1024, C = 4096;
int N = B * T;
size_t bytes = N * C * sizeof(float);
size_t weights_bytes = C * sizeof(float);
float *h_data = (float *)malloc(bytes);
float *h_out = (float *)malloc(bytes);
float *out = (float *)malloc(bytes);
float *h_weights = (float *)malloc(weights_bytes);

init(h_data, N * C, -1.0, 1.0);
init(h_weights, C, 0, 1.0);

rmsnorm_cpu(h_out, h_data, h_weights, N, C);

float *d_data, *d_out, *d_weights;
cudaMalloc(&d_data, bytes);
cudaMalloc(&d_out, bytes);
cudaMalloc(&d_weights, weights_bytes);
cudaMemcpy(d_data, h_data, bytes, cudaMemcpyHostToDevice);
cudaMemcpy(d_weights, h_weights, weights_bytes, cudaMemcpyHostToDevice);

kernel<<<N, BLOCK_SIZE>>>(d_out, d_data, d_weights, N, C);
cudaMemcpy(out, d_out, bytes, cudaMemcpyDeviceToHost);

check_results(h_out, out, N * C, 1e-5);

free(h_data);
free(h_weights);
free(h_out);
free(out);
cudaFree(d_data);
cudaFree(d_out);
cudaFree(d_weights);

return 0;
}

5. CUDA版本V2一个Block处理一行

当数据量比较大的时候,一个Warp处理一行,显然是不够的,这个时候,我们就可以使用一个Block处理一行,总共需要B*T*Block_size个线程,我们仍然将其过程分为三个阶段。

  1. 计算均方根。这里我们就不能使用__shfl_down_sync等洗牌函数了,因为这个函数是Warp间的规约,不适用于Block,所以,我们只能引入共享内存。使用共享内存的方法已经在第一节CUDA算子优化(1):Reduce中讲过,请参考第一节。
  2. 计算RMS(x)
  3. 乘上缩放因子,计算最终结果。
#include <stdio.h>
#include <stdlib.h>

#define BLOCK_SIZE 32

__global__ void kernel(float *d_out, float *d_data, float *weights, int N, int C) {
int index = threadIdx.x + blockDim.x * blockIdx.x;
int row = index / blockDim.x;
int tid = index % blockDim.x;
if (row < N) {
float *inp = d_data + row * C;
float *out = d_out + row * C;

int factor = C / BLOCK_SIZE;
__shared__ float smem[BLOCK_SIZE];
float local_sum = 0.0f;
for (int i = 0; i < factor; ++i) {
local_sum += inp[tid + i * BLOCK_SIZE] * inp[tid + i * BLOCK_SIZE];
}
smem[tid] = local_sum;
__syncthreads();
for (int stride = BLOCK_SIZE / 2; stride >= 1; stride /= 2) {
if (tid < stride) {
smem[tid] += smem[tid + stride];
}
__syncthreads();
}
float norm = 1/ sqrt(smem[0] / C + 1e-6);

for (int i = 0; i < C; ++i) {
out[i] = inp[i] * norm * weights[i];
}
}
}

void init(float *data, int N, int low, int high) {
for (int i = 0; i < N; ++i) {
float r = (float)rand() / RAND_MAX;
data[i] = low + (high - low) * r;
}
}

void rmsnorm_cpu(float *h_out, float *h_data, float *weights, int N, int C) {
for (int i = 0; i < N; ++i) {
float *inp = h_data + i * C;
float *out = h_out + i * C;

float local_sum = 0.0f;
for (int j = 0; j < C; ++j) {
local_sum += inp[j] * inp[j];
}
local_sum = 1 / sqrt(local_sum / C + 1e-6);
for (int j = 0; j < C; ++j) {
out[j] = inp[j] * local_sum * weights[j];
}
}
}

int check_results(const float *actual, const float *expected,
int n, float tol) {
int errors = 0;
float max_diff = 0.0f;
for (int i = 0; i < n; i++) {
float diff = fabsf(actual[i] - expected[i]);
if (diff > max_diff)
max_diff = diff;
if (diff > tol)
errors++;
}
printf(" max_diff = %.2e\n", max_diff);
printf(" GPU vs CPU : %s\n", errors == 0 ? "PASS" : "FAIL");
return errors == 0;
}

int main() {
int B = 64, T = 1024, C = 2048;
int N = B * T;
size_t bytes = B * T * C * sizeof(float);
size_t weights_bytes = C * sizeof(float);
float *h_data = (float *)malloc(bytes);
float *h_weights = (float *)malloc(weights_bytes);
float *h_out = (float *)malloc(bytes);
float *out = (float *)malloc(bytes);

init(h_data, B * T * C, -1.0, 1.0);
init(h_weights, C, 0, 1.0);
rmsnorm_cpu(h_out, h_data, h_weights, B * T, C);

float *d_data, *d_out, *d_weights;
cudaMalloc(&d_data, bytes);
cudaMalloc(&d_out, bytes);
cudaMalloc(&d_weights, weights_bytes);
cudaMemcpy(d_data, h_data, bytes, cudaMemcpyHostToDevice);
cudaMemcpy(d_weights, h_weights, weights_bytes, cudaMemcpyHostToDevice);

kernel<<<N, BLOCK_SIZE>>>(d_out, d_data, d_weights, N, C);
cudaMemcpy(out, d_out, bytes, cudaMemcpyDeviceToHost);

check_results(h_out, out, N * C, 1e-5);

free(h_data);
free(h_weights);
free(h_out);
free(out);
cudaFree(d_data);
cudaFree(d_out);
cudaFree(d_weights);

return 0;
}