LeetGPU习题07:Norm系列代码实现

6854 字
34 分钟
LeetGPU习题07:Norm系列代码实现
2026-06-12

Pytorch调用#

PyTorch的调用接口是高度统一的,它们都继承自nn.Module,核心都是沿着某个维度计算统计量,之后标准化,可选仿射变换。

BN、LN、RMS的区别

区别在于沿着哪个轴做归一化,以及计算哪些统计量

理解了这个本质,剩下的就是查找参数,看维度。

BatchNorm#

torch.nn.BatchNorm1d(num_features, eps=1e-05, momentum=0.1, affine=True, ...)
torch.nn.BatchNorm2d(num_features, eps=1e-05, momentum=0.1, affine=True, ...)
torch.nn.BatchNorm3d(num_features, eps=1e-05, momentum=0.1, affine=True, ...)

对于BatchNorm来说,主要有上述三个版本,区别在于它们对于输入的维度布局不同,我们来看一下基本用法:

def batchnorm_example():
x = torch.rand(100, 16, 784)
layer = nn.BatchNorm1d(16)
out = layer(x)
y = torch.rand(1, 16, 7, 7)
layer = nn.BatchNorm2d(16) # 传入通道数C
out = layer(y)

Pytorch要求显式传入特征通道数,以便可以在内部初始化可学习参数weight和bias。它们的形状都是(C,), 实际上BatchNorm会根据自身是否为training来自动选择后续行为,对于调用者是完全透明的。

学习了官方的调用之后,我们来看一下具体是如何实现的:

def batchnorm(x, running_mean, running_var, weight, bias, training, momentum=0.1, eps=1e-5):
"""
x.dim() 返回张量的维度数,例如对于一个形状为 (3, 4, 5) 的张量,x.dim() 将返回 3。
range(x.dim()) 生成一个从 0 到 x.dim() - 1 的序列,例如生成 [0, 1, 2]。
[d for d in range(x.dim()) if d != 1] 是列表推导式,它遍历上述生成的序列,并将不等于 1 的维度索引添加到列表中。
"""
dims = [d for d in range(x.dim()) if d != 1]
"""
[1: -1]是一个列表,第一个元素是1,第二个元素是-1
[1] * (x.dim() - 2) 是一个列表,包含 x.dim() - 2 个元素,每个元素都是1。
shape表示前两个维度是1和-1,后边的维度全都是1
"""
shape = [1, -1] + [1] * (x.dim() - 2)
# 训练模式
if training:
mean = x.mean(dim=dims, keepdim=True) # shape: (1, C, 1, 1)
var_biased = x.var(dim=dims, keepdim=True, correction=0) # 有偏估计,用于当前批次的归一化
var_unbiased = x.var(dim=dims, keepdim=True, correction=1) # 无偏估计,用于更新全局统计量
with torch.no_grad():
# running_mean 和running_var 形状是(C,), 用squeeze()挤掉多余的维度
running_mean.data = (1 - momentum) * running_mean.data + momentum * mean.squeeze()
running_var.data = (1 - momentum) * running_var.data + momentum * var_unbiased.squeeze()
var = var_biased
else:
mean = running_mean.view(shape)
var = running_var.view(shape)
# 归一化
x_norm = (x - mean) / torch.sqrt(var + eps)
# 仿射变换:y=γx+β
if weight is not None:
w = weight.view(shape)
b = bias.view(shape)
return x_norm * w + b
return x_norm

具体内容如上,关键点我都已经写出来了,大家自行查阅即可。

优化思维时时刻刻都要有,我还写了一个优化版本的,mean、var_unbiased可以同时计算出来:

def manual_batchnorm_v2(x, running_mean, running_var, weight, bias, training, momentum=0.1, eps=1e-5):
dims = [d for d in range(x.dim()) if d != 1]
shape = [1, -1] + [1] * (x.dim() - 2)
if training:
# 计算参与归约的元素总数
n = 1
for d in dims:
n *= x.shape[d]
# 一次性计算均值和方差(无偏方差)
var_unbiased, mean = torch.var_mean(x, dim=dims, keepdim=True, correction=1)
# 转为有偏方差用于归一化
var_biased = var_unbiased * ((n - 1) / n)
with torch.no_grad():
running_mean.data = (1 - momentum) * running_mean.data + momentum * mean.squeeze()
running_var.data = (1 - momentum) * running_var.data + momentum * var_unbiased.squeeze()
var = var_biased
else:
mean = running_mean.view(shape)
var = running_var.view(shape)
x_norm = (x - mean) / torch.sqrt(var + eps)
if weight is not None:
return x_norm * weight.view(shape) + bias.view(shape)
return x_norm

LayerNorm#

LayerNorm不依赖batch统计量,在NLP和Transformer是绝对主力。

torch.nn.LayerNorm(normalized_shape, eps=1e-05, elementwise_affine=True, ...)
# shape可以是int或者tuple,表示输入最后几个维度的大小,归一化在这些维度进行
# eps,同BatchNorm
# elementwise_affine,是否学习逐元素的weight和bias

调用示例如下:

ln = nn.LayerNorm(512)
x = torch.randn(32, 10, 512)
out = ln.(x) # 形状不变,每个token的512维向量被归一化

我们手动实现的版本:

import torch
import torch.nn as nn
class layernorm(nn.Module):
def __init__(self, embed_size, eps=1e-5):
super().__init__()
self.gamma = nn.Parameter(torch.ones(embed_size))
self.beta = nn.Parameter(torch.zeros(embed_size))
self.eps = eps
def forward(self, x):
# correction=0表示有偏估计
var, mean = torch.var_mean(x, dim=-1, keepdim=True, correction=0)
x_norm = (x - mean) / torch.sqrt(var + self.eps)
return x_norm * self.gamma + self.beta
if __name__ == "__main__":
x = torch.randn(2, 4, 3)
embed_size = x.size(-1)
official = nn.LayerNorm(embed_size, eps=1e-5)
my = layernorm(embed_size, eps=1e-5)
# 复制相同权重
my.gamma.data = official.weight.data.clone()
my.beta.data = official.bias.data.clone()
diff = (official(x) - my(x)).abs().max().item()
print(f"Max difference: {diff:.2e}")

RMSNorm#

RMSNorm是LayerNorm的一个简化变体,直接减去均值计算的步骤,只除以均方根,在最新的大模型里广泛使用。

torch.nn.RMSNorm(normalized_shape, eps=1e-06, elementwise_affine=True, ...)

使用方法完全同LayerNorm一致。

特性BatchNormLayerNormRMSNorm
归一化轴沿 batch 和空间维度 (每个通道独立)沿特征维度 (每个样本独立)沿特征维度 (每个样本独立)
依赖 batch
可学习参数weight + bias (通道数)weight + bias (归一化形状)weight
适用场景CNN、CVNLP、Transformer大模型
PyTorch 接口BatchNorm1d/2d/3dLayerNormRMSNorm

Triton实现#

在了解了Pytorch这一层级的实现之后,我们要手动在Triton中实现LN和RMS。

下边我们核心先讲解LN和RMS的Triton版本,BN我就先忽略,大家可以自行学习。

归一化的视角#

无论是LN还是RMS,输入都是一个二维矩阵XRN×KX \in R^{N \times K},其中N是所有token的总行数,K是特征维度。

每一行的计算是完全独立的。 对于第n行,算子聚合该行内所有K个元素,算出指定的统计量,再用统计量去逐元素变换该行。

Tip

这是经典的Reduce+Broadcast模式

具体来说,LayerNorm需要两步聚合,分别计算均值和方差;RMSNorm只进行一次平方和聚类。

在GPU上,Reduce是代价最高的部分,Broadcast是element wise操作。很明显,Norm操作是带宽密集型操作。

LayerNorm#

我在这里参考了flash-attenion官方的实现: flash-attention

关于反向传播的代码

实际上,除了前向传播,在写绝大多数算子的时候,也要考虑反向传播,后续我会写一个完整的macroTorch,来完整的学习是如何进行前向传播与反向传播的。

关于自动调优,可以写一个函数来囊括进去常用的内容。 同时我在这里保存了mean和rstd,这样能够给未来的反向传播直接复用,避免重新进行规约计算,因此可以看到我显式的进行了两次保存。

并且还有一个设计上的巧思,即步长参数化,而不是假设内存严格连续,这可以让kernel实现更加复杂的张量布局。

import torch
import torch.nn as nn
import triton
import triton.language as tl
from triton.testing import do_bench
def get_autotune_configs():
warp_size = 32
max_threads_per_block = 1024
configs = []
for num_warps in [1,2, 4, 8, 16, 32]:
if num_warps * warp_size <= max_threads_per_block:
configs.append(triton.Config({}, num_warps=num_warps))
return configs
@triton.autotune(
configs=get_autotune_configs(),
key=["N"],
)
@triton.jit
def layer_norm_fwd_kernel(
X_ptr, Y_ptr, W_ptr, B_ptr, Mean_ptr, Rstd_ptr,
stride_x_row, stride_y_row, # 传入行步长即可灵活索引,Triton编程常见模式
N, eps, BLOCK_N: tl.constexpr
):
# 每个program处理1行
row_idx = tl.program_id(0)
X_row_ptr = X_ptr + row_idx * stride_x_row
Y_row_ptr = Y_ptr + row_idx * stride_y_row
cols = tl.arange(0, BLOCK_N)
mask = cols < N
# 加载该行的所有元素,在计算中提升为FP32
x = tl.load(X_row_ptr + cols, mask=mask, other=0.0).to(tl.float32)
w = tl.load(W_ptr + cols, mask=mask).to(tl.float32)
b = tl.load(B_ptr + cols, mask=mask).to(tl.float32)
# 规约:均值 and 方差
mean = tl.sum(x, axis=0) / N
tl.store(Mean_ptr + row_idx, mean)
x_bar = tl.where(mask, x - mean, 0.0)
var = tl.sum(x_bar * x_bar, axis=0) / N
rstd = 1.0 / tl.sqrt(var + eps)
tl.store(Rstd_ptr + row_idx, rstd)
# 广播:归一化+仿射变换
y = (x - mean) * rstd * w + b
tl.store(Y_row_ptr + cols, y, mask=mask)
def layer_norm_fwd(x, weight, bias=None, eps=1e-5):
M, N = x.shape
y = torch.empty_like(x)
mean = torch.empty(M, device=x.device, dtype=torch.float32)
rstd = torch.empty(M, device=x.device, dtype=torch.float32)
if bias is None:
bias = torch.zeros(N, device=x.device, dtype=weight.dtype)
if weight is None:
weight = torch.ones(N, device=x.device, dtype=x.dtype)
BLOCK_N = triton.next_power_of_2(N)
layer_norm_fwd_kernel[(M,)](
x, y, weight, bias, mean, rstd,
x.stride(0), y.stride(0),
N, eps, BLOCK_N=BLOCK_N
)
return y, mean, rstd
def test_correctness(shapes=[(128, 256), (512, 1024)]):
for M, N in shapes:
x = torch.randn(M, N, device='cuda', dtype=torch.float32)
weight = torch.randn(N, device='cuda', dtype=torch.float32)
bias = torch.randn(N, device='cuda', dtype=torch.float32)
eps = 1e-5
ln = nn.LayerNorm(N, eps=eps).to('cuda')
ln.weight.data = weight
ln.bias.data = bias
y_ref = ln(x)
y_tri, _, _ = layer_norm_fwd(x, weight, bias, eps)
max_diff = (y_tri - y_ref).abs().max().item()
print(f"Shape ({M}, {N}): max diff = {max_diff:.6e}")
bench_perf_report = None
try:
from triton.testing import Benchmark, perf_report
@perf_report(
Benchmark(
x_names=["N"],
x_vals=[256, 512, 1024, 2048, 4096, 8192],
line_arg="provider",
line_vals=["triton", "pytorch"],
line_names=["Triton", "PyTorch"],
styles=[("blue", "-"), ("red", "-")],
ylabel="Latency (ms)",
plot_name="LayerNorm Fwd Performance",
args={"M": 1024, "eps": 1e-5, "dtype": torch.float32},
)
)
def _bench_perf_report(M, N, eps, dtype, provider):
device = 'cuda'
x = torch.randn(M, N, device=device, dtype=dtype)
weight = torch.randn(N, device=device, dtype=dtype)
bias = torch.randn(N, device=device, dtype=dtype)
if provider == "triton":
def run():
return layer_norm_fwd(x, weight, bias, eps)
else:
ln = nn.LayerNorm(N, eps=eps).to(device)
ln.weight.data = weight
ln.bias.data = bias
def run():
return ln(x)
return do_bench(run, quantiles=[0.5, 0.2, 0.8])
bench_perf_report = _bench_perf_report
except ImportError:
print("当前 Triton 版本不支持 perf_report,跳过高阶绘图功能\n")
if __name__ == "__main__":
test_correctness()
if bench_perf_report is not None:
bench_perf_report.run(show_plots=True, print_data=True, save_path="./layer_norm_fwd_perf.png")

性能对比如下:

LN在Triton和Pytorch的性能表现
LN在Triton和Pytorch的性能表现

RMSNorm#

LayerNorm需要计算两个统计量:mean和std。 mean就得涉及到规约,而RMSNorm只计算一个:

rms = sqrt(E[x^2] + eps)
y = x / rms * γ

核心代码如下:

@triton.autotune(
configs=autotune_configs(),
key=["N"],
)
@triton.jit
def rms_norm_fwd_kernel(
X_ptr, Y_ptr, W_ptr,
stride_x_row, stride_y_row,
N, eps,
BLOCK_N: tl.constexpr,
):
row_idx = tl.program_id(0)
X_row_ptr = X_ptr + row_idx * stride_x_row
Y_row_ptr = Y_ptr + row_idx * stride_y_row
cols = tl.arange(0, BLOCK_N)
mask = cols < N
x = tl.load(X_row_ptr + cols, mask=mask, other=0.0).to(tl.float32)
w = tl.load(W_ptr + cols, mask=mask, other=0.0).to(tl.float32)
# RMSNorm: rstd = rsqrt(mean(x^2) + eps)
x2 = x * x
mean_x2 = tl.sum(x2, axis=0) / N
rstd = tl.rsqrt(mean_x2 + eps)
# 广播:缩放+仿射变换
y = x * rstd * w
tl.store(Y_row_ptr + cols, y, mask=mask)
def rms_norm_fwd(x, weight, eps=1e-5):
M, N = x.shape
y = torch.empty_like(x)
if weight is None:
weight = torch.ones(N, device=x.device, dtype=x.dtype)
BLOCK_N = triton.next_power_of_2(N)
rms_norm_fwd_kernel[(M,)](
x, y, weight,
x.stride(0), y.stride(0),
N, eps,
BLOCK_N=BLOCK_N,
)
return y

RMS的Triton与Pytorch性能对比图
RMS的Triton与Pytorch性能对比图
我在特征维度从256到8192进行了测试,Triton版本始终优于Pytorch,并且自动调优保证了在不同尺寸下的均衡表现。

RMS在FP16下的优势更加明显,因为Triton对混合精度路径的控制更加精细,避免了不必要的类型转换开销。

====================================================================================================
RMSNorm Performance: dtype=torch.float16
====================================================================================================
M N Triton(ms) PyTorch (ms) my vs PT
----------------------------------------------------------------------------------------------------
128 256 0.003816 0.021189 5.55x
128 512 0.004489 0.021973 4.89x
128 1024 0.004901 0.023449 4.78x
128 2048 0.005038 0.026430 5.25x
128 4096 0.006675 0.032510 4.87x
128 8192 0.008356 0.038082 4.56x
512 256 0.004816 0.022823 4.74x
512 512 0.005085 0.025100 4.94x
512 1024 0.006658 0.029662 4.45x
512 2048 0.008573 0.038389 4.48x
512 4096 0.011734 0.058861 5.02x
512 8192 0.022541 0.093368 4.14x
1024 256 0.005728 0.025211 4.40x
1024 512 0.006578 0.029643 4.51x
1024 1024 0.009162 0.037848 4.13x
1024 2048 0.011643 0.057440 4.93x
1024 4096 0.020953 0.093154 4.45x
1024 8192 0.041657 0.170544 4.09x
2048 256 0.006578 0.029480 4.48x
2048 512 0.007969 0.037949 4.76x
2048 1024 0.011659 0.057081 4.90x
2048 2048 0.021078 0.092297 4.38x
2048 4096 0.041632 0.169116 4.06x
2048 8192 0.077019 0.681771 8.85x
4096 256 0.007757 0.037739 4.87x
4096 512 0.012584 0.057165 4.54x
4096 1024 0.020957 0.091967 4.39x
4096 2048 0.040349 0.168362 4.17x
4096 4096 0.077127 0.669127 8.68x
4096 8192 0.149073 1.469917 9.86x
====================================================================================================
RMSNorm Performance: dtype=torch.float32
====================================================================================================
M N Triton(ms) PyTorch (ms) my vs PT
----------------------------------------------------------------------------------------------------
128 256 0.004033 0.016027 3.97x
128 512 0.004791 0.016718 3.49x
128 1024 0.005044 0.018569 3.68x
128 2048 0.006306 0.021332 3.38x
128 4096 0.009204 0.026772 2.91x
128 8192 0.013090 0.030454 2.33x
512 256 0.005348 0.017874 3.34x
512 512 0.006673 0.019975 2.99x
512 1024 0.008565 0.023732 2.77x
512 2048 0.012434 0.031180 2.51x
512 4096 0.020615 0.048329 2.34x
512 8192 0.043017 0.077455 1.80x
1024 256 0.006282 0.019679 3.13x
1024 512 0.008333 0.023323 2.80x
1024 1024 0.012538 0.030153 2.40x
1024 2048 0.020061 0.046295 2.31x
1024 4096 0.043250 0.076708 1.77x
1024 8192 0.079363 0.145954 1.84x
2048 256 0.008264 0.023329 2.82x
2048 512 0.011833 0.030187 2.55x
2048 1024 0.023103 0.046089 1.99x
2048 2048 0.042146 0.075980 1.80x
2048 4096 0.078298 0.144876 1.85x
2048 8192 0.150213 0.479396 3.19x
4096 256 0.012040 0.030257 2.51x
4096 512 0.020640 0.045097 2.18x
4096 1024 0.040291 0.075542 1.87x
4096 2048 0.079232 0.144678 1.83x
4096 4096 0.150072 0.465174 3.10x
4096 8192 0.296848 1.028277 3.46x

FP32下加速比在2-4倍之间,FP16下可达4-9倍。

CUDA实现#

二者对于每一行N,都会遍历该行的所有元素,计算统计量(平方和、均值和方差等),并且使用这些统计量做归一化变换。

实际上,Norm的核心分为两个阶段:

阶段 1 — Reduce(归约):
每一行的 K 个元素 → 汇总成 1 个统计量(标量)
例如:RMSNorm 把一行 8192 个 float 聚合成一个平方和
LayerNorm 把一行 8192 个 float 先聚合成均值,再聚合成方差
阶段 2 — Broadcast(广播):
用这 1 个标量去变换该行的每一个元素
y[n][k] = f(x[n][k], stat[n])

用 C++ 伪代码表达就是:

for (int n = 0; n < N; n++) { // N 行,行间独立
float stat = 0;
for (int k = 0; k < K; k++) { // Reduce: K 个元素 → 1 个值
stat += compute(x[n][k]);
}
stat = finalize(stat); // 例如 rsqrt(stat/K + eps)
for (int k = 0; k < K; k++) { // Broadcast: 1 个值 → K 个元素
y[n][k] = transform(x[n][k], stat, gamma[k], beta[k]);
}
}

为什么Norm是带宽瓶颈#

y = x * rsqrt(mean(x²) + ε) * gamma

每读取一个元素只做两次乘加和一次 rsqrt,算术密度低。 这类 kernel 的性能不由算力决定,而由内存带宽决定。

对于4090D,理论带宽为1000GB/s左右,一个4096×81924096 \times 8192的FP16的RMSNorm理论最优耗时为:

数据量 = (读x + 读gamma + 写y) = (2×4096×8192 + 8192) × 2 bytes ≈ 134 MB
理论最优 = 134 MB / 1008 GB/s ≈ 0.133 ms

测试结果如下:

❯ ./rms_norm
--- RMSNorm 性能测试 ---
kernel | shape | avg_ms | bandwidth | correctness
--------------------------------------------------------------------
RMSNorm FP32 | 1024x256 | 0.007 ms | 310.1 GB/s | err: 6.0e-07 [OK]
RMSNorm FP16 | 1024x256 | 0.007 ms | 156.1 GB/s | err: 1.5e-03 [OK]
RMSNorm FP32 | 1024x1024 | 0.010 ms | 876.1 GB/s | err: 9.5e-07 [OK]
RMSNorm FP16 | 1024x1024 | 0.008 ms | 556.0 GB/s | err: 1.5e-03 [OK]
RMSNorm FP32 | 2048x4096 | 0.078 ms | 865.7 GB/s | err: 2.0e-06 [OK]
RMSNorm FP16 | 2048x4096 | 0.037 ms | 905.3 GB/s | err: 1.6e-03 [OK]
RMSNorm FP32 | 4096x8192 | 0.294 ms | 914.2 GB/s | err: 4.4e-06 [OK]
RMSNorm FP16 | 4096x8192 | 0.147 ms | 911.0 GB/s | err: 1.6e-03 [OK]

大概在理论性能的90%左右,基本能够跑满带宽。

Oneflow的设计分析#

OneFlow 的 LayerNorm/RMSNorm 实现是目前开源社区中架构最成熟的实现之一。 在这里我还是参考它的设计思路。

分离策略#

1. 数据搬运与计算分离

// OneFlow 将读写抽象为仿函数
template<typename SRC, typename DST>
struct DirectLoad {
template<int N>
__device__ void load(DST* dst, int64_t row, int64_t col) const { ... }
};

同一个Kernel模板可以适配不同的输入输出需求,Kernel只关心”加载 → 计算 → 存回”,不关心数据从哪里来,到哪里去。

2. 规约与归一化分离

朴素实现中,规约和归一化交织在一起,例如需要先完整遍历一次数据求出 μ\mu,再遍历第二次求方差,第三次才做归一化。这不仅 I/O 开销大,而且求方差时容易产生严重的数值精度问题。

朴素求方差的公式是:

σ2=E[x2](E[x])2\sigma^2 = E[x^2] - (E[x])^2

这要求同时累积 x\sum xx2\sum x^2。当数据的均值较大而方差本身很小时,x2\sum x^2(x)2/N(\sum x)^2/N 是两个非常接近的大数,相减会导致大量有效数字抵消,这叫做 catastrophic cancellation,结果误差极大。

Welford 算法提供了一种数值稳定的在线计算方式,它不需要保存 x2\sum x^2,而是维护一个与均值无关的修正平方和m2

template<typename T>
__device__ void WelfordCombine(T val, T* mean, T* m2, T* count) {
*count += 1;
T delta1 = val - *mean;
*mean += delta1 / *count;
T delta2 = val - *mean;
*m2 += delta1 * delta2;
}

Welford 算法天然支持单趟遍历,一边读取数据一边更新 meanm2

GPU 上求全局均值/方差需要跨线程块规约。Welford 具备可合并性 这允许我们先让每个 warp / block 独立做局部 Welford,然后只归并这几个状态.

一旦得到全局 μ\muσ\sigma,归一化就变成了一个完全无依赖的逐元素操作(x - μ) * inv_std
这可以很方便地和其他逐元素操作(如 γ,β\gamma, \beta、残差加法、激活函数)融合到一个 kernel 里,形成端到端的“大融合算子”。 分离的设计让规约成为一次全局通信+少量计算,归一化成为纯本地并行计算,边界清晰,利于编译器和手写 kernel 深度优化。

3. 策略与算法分离

策略适用条件线程组织核心思路
WarpImplK ≤ 10242D block一个 warp group 一行,block 内多行并行
BlockSMemImplK > 10241D block + dynamic shared memory数据缓存到 shared memory,省去第二次 global 读
BlockUncachedImplSMemImpl 放不下1D block,两次 global 读依赖 L2 cache

设计策略#

1. 基于 occupancy 的动态 grid sizing

cudaOccupancyMaxActiveBlocksPerMultiprocessor(&max_active_blocks, func, block_size, smem);
*num_blocks = max(1, min(max_blocks, sm_count * max_active_blocks * waves));

根据 SM 数量和每个 SM 能驻留的最大 block 数,乘以一个 waves 因子(默认 32),确保有足够的 block 来隐藏内存延迟。可以根据当前占用动态设计grid size大小。

2. Pack 与 pack_size 自动选择

if (cols % 4 == 0 && CanPackAs<LOAD>(load, 4)) {
// use pack_size = 4
} else if (cols % 2 == 0 && CanPackAs<LOAD>(load, 2)) {
// use pack_size = 2
} else {
// fallback to 1
}

根据列对齐自动选择最优向量化宽度,在不牺牲通用性的前提下最大化访存效率。

3. 小K大N场景的批量处理

当 K 很小(如 64)且 N 很大(如 10⁶)时,每个 warp group 一次处理 2 行而非 1 行,指令级并行度翻倍,更好地隐藏计算延迟。

我的CUDA设计#

基于对 OneFlow 的深入分析,我进行了精简,保留代码的核心思路。

模块架构预览#

cuda_norm/
├── reduce.cuh ← warp/block 级规约原语
├── io.cuh ← 向量化访存抽象 (Pack, DirectLoad, AffineStore)
├── norm_kernel.cuh ← 通用 kernel 骨架 (WarpImpl + BlockSMemImpl)
├── rms_norm.cuh ← RMSNorm: Stats特化 + dispatch + host API
├── layer_norm.cuh ← LayerNorm: Stats特化 + dispatch + host API
└── bench/
├── reduce_bench.cu
├── io_bench.cu
├── rms_norm_bench.cu
└── layer_norm_bench.cu

依赖关系如下:

reduce.cuh → io.cuh → norm_kernel.cuh → rms_norm.cuh
→ layer_norm.cuh

Reduce模块#

职责

提供Warp级别和Block级别的求和规约。

核心接口如下

template <typename T>
__device__ T warp_reduce_sum(T val);
template <const int NUM_THREADS, typename T>
__device__ T block_reduce_sum(T val);

设计思路如下

Warp内部使用__shfl_xor_sync()

初始: lane0 lane1 lane2 lane3 ... lane31
mask=16: 每个 lane 与相距16的 lane 交换并求和
mask=8: 每个 lane 与相距8的 lane 交换并求和
...
mask=1: 最终每个 lane 持有全部32个值的和

Block规约是两级结构:

第一步: 各 warp 内部独立规约 → lane0 写入 shared memory[warp_id]
第二步: warp0 从 shared memory 读取各 warp 结果 → 再做一次 warp 规约

具体代码可以在代码仓库中查看。

IO模块#

职责

将数据搬运与计算逻辑解耦。提供向量化访存(128-bit LDG/STG)、透明的类型转换(fp16↔fp32)、 以及零开销的 affine 融合(gamma/beta 乘加),让 kernel 代码不感知内存布局、对齐和精度。

核心接口如下

// 层0: 16字节对齐的寄存器数组,触发 LDG.128 / STG.128
template <typename T, int N>
struct alignas(sizeof(T) * N) Pack { T elem[N]; };
// 层1: 加载 + 类型转换合二为一 (如 half→float)
template <typename SRC, typename DST>
struct DirectLoad {
template <int N>
__device__ void load(DST* dst, int64_t row, int64_t col) const;
};
// 层2: 类型转换 + 向量化存储
template <typename SRC, typename DST>
struct DirectStore {
template <int N>
__device__ void store(const SRC* src, int64_t row, int64_t col);
};
// 层3: 存储时融合 affine 参数 (y = x * gamma + beta),编译期分支消除
template <typename SRC, typename DST, bool do_scale, bool do_center>
struct AffineStore {
template <int N>
__device__ void store(const SRC* src, int64_t row, int64_t col);
};

具体代码可以在代码仓库中查看。

假设没有IO层,直接在Kernel中写数据搬运代码,是什么情况?

template <int N>
__global__ void rms_norm_naive(const __half* x, __half* y, const __half* gamma,
int rows, int cols, float eps) {
int row = blockIdx.x;
int K = cols;
// 问题1: 逐元素标量加载 — 32 条 LDG.16 指令,而非 4 条 LDG.128
float sum_sq = 0.0f;
for (int k = threadIdx.x; k < K; k += blockDim.x) {
float v = __half2float(x[row * K + k]); // 问题2: 类型转换散落各处
sum_sq += v * v;
}
// 问题3: reduce 逻辑和访存耦合在一起
// ...
float inv_rms = rsqrtf(sum_sq / K + eps);
// 问题4: 写回时再次手动转换 + 手动乘 gamma, 每次都是标量 STG.16
for (int k = threadIdx.x; k < K; k += blockDim.x) {
float v = __half2float(x[row * K + k]);
float out = v * inv_rms;
// 问题5: gamma 的加载和乘法也在 kernel 中,换了 LayerNorm 要全部重写
float g = __half2float(gamma[k]);
y[row * K + k] = __float2half(out * g);
}
}

有了IO层之后是什么样子的呢?

template <typename LOAD, typename STORE, typename ComputeType, typename Stats, ...>
__global__ void NormWarpImpl(LOAD load, STORE store, int rows, int cols, float eps) {
ComputeType buf[cols_per_thread]; // 寄存器中的缓冲区
// 加载: 一行代码,封装了 128-bit LDG + half→float 转换
load.template load<pack_size>(buf, row, col);
// 计算: kernel 只看到 float 类型的数据,不需要知道原始是 fp16 还是 fp32
Stats::normalize(buf, stat, cols_per_thread);
// 存储: 封装了 float→half 转换 + gamma*out + 128-bit STG
store.template store<pack_size>(buf, row, col);

编译器辅助向量化#

template <typename T, int N>
struct alignas(sizeof(T) * N) Pack {
T elem[N];
};

CUDA 编译器在看到对齐的 16 字节读写时,会自动生成 LDG.128/STG.128 指令。

alignas是核心
Pack<float, 4>: 4 × 4 bytes = 16 bytes, alignas(16)
Pack<__half, 8>: 8 × 2 bytes = 16 bytes, alignas(16)

在这里还有一个是T elem[N],这里Pack放在了寄存器,这是因为每个线程处理的元素比较少,并且没有访存冲突和数据同步。

类型转换透明化#

template <typename SRC, typename DST>
struct DirectLoad {
template <int N>
__device__ void load(DST* dst, int64_t row, int64_t col) const {
Pack<SRC, N> pack;
const int64_t offset = (row * row_size + col) / N;
pack = *reinterpret_cast<const Pack<SRC, N>*>(src + offset * N); // 128-bit LDG
#pragma unroll
for (int i = 0; i < N; ++i)
dst[i] = static_cast<DST>(pack.elem[i]); // 自动类型转换
}
};

SRCDST 两个模板参数是这一层的核心设计

  • SRC = 显存中的存储类型(__half 用于 FP16 模型)
  • DST = 寄存器中的计算类型(float 用于 FP32 精度计算)
GPU 显存: fp16 ──DirectLoad<__half, float>──→ 寄存器: fp32
自动调用 __half2float
那这里为什么不让Kernel手动类型转换呢?

当然是,容易忘!

DirectStoreDirectLoad 的镜像,方向相反,逻辑完全对称:

pack = *reinterpret_cast<const Pack<SRC, N>*>(src + offset * N); // Load: Global → Pack
*reinterpret_cast<Pack<DST, N>*>(dst + offset * N) = pack; // Store: Pack → Global

AffineStore编译期分支消除#

template <typename SRC, typename DST, bool do_scale, bool do_center>
struct AffineStore {
template <int N>
__device__ void store(const SRC* src, int64_t row, int64_t col) {
if (do_scale) // 编译期 false → 整块代码被删除
gamma_pack = *reinterpret_cast<const Pack<DST, N>*>(gamma + w_offset * N);
if (do_center) // 编译期 false → 整块代码被删除
beta_pack = *reinterpret_cast<const Pack<DST, N>*>(beta + w_offset * N);
for (int i = 0; i < N; ++i) {
DST v = static_cast<DST>(src[i]);
if (do_scale) v = v * gamma_pack.elem[i]; // RMSNorm: 保留 | 纯Norm: 删除
if (do_center) v = v + beta_pack.elem[i]; // LayerNorm: 保留 | RMSNorm: 删除
dst_pack.elem[i] = v;
}
*reinterpret_cast<Pack<DST, N>*>(dst + offset * N) = dst_pack; // 128-bit STG
}
};

我在这里写到了两个编译期参数,这样可以只需要传入true或者false,就能够分别在RMS和LN中进行调用。

同时,gamma/beta 的加载也只在需要时才发生

if (do_scale)
gamma_pack = *reinterpret_cast<const Pack<DST, N>*>(gamma + w_offset * N);

do_scale=false 时,gamma 的加载指令都不会生成。这意味着纯归一化场景下,gamma 可以传 nullptr 而不会崩溃。

完整的数据流#

把三层抽象串联起来,一次 Norm kernel 调用的数据通路由 4 步完成:

GPU Global Memory
┌──────────────────┐
│ x (fp16) │
│ gamma (fp16) │
│ beta (fp16) │
│ y (fp16) │
└──┬────────────┬──┘
│ ▲
┌──────────────┘ └──────────────┐
│ ① DirectLoad │ ③ AffineStore
│ LDG.128 × N/pack_size │ LDG.128 (gamma/beta)
│ half→float 转换 │ float→half 转换
│ │ val * gamma + beta
▼ │ STG.128
┌─────────┐ │
│ buf[] │ ② Stats::reduce + normalize │
│ (寄存器) │──────────→ stat ──→ normalize ──────→┘
└─────────┘

使用示例#

RMSNorm (乘 gamma,不加 beta)

// 准备数据
const __half* x = ...; // 输入 [N, K]
__half* y = ...; // 输出 [N, K]
const __half* gamma = ...; // 权重 [K]
// 创建 Loader 和 Store
DirectLoad<__half, float> load(x, K);
AffineStore<float, __half, true, false> store(y, K, gamma, nullptr);
// ^^^^ ^^^^^
// do_scale=true do_center=false
// 在 kernel 中使用
float buf[8];
load.load<8>(buf, row, col); // 加载 8 个 half → float
// ... 归一化计算 (buf 中是 float) ...
store.store<8>(buf, row, col); // 存储: float→half, 乘 gamma

LayerNorm (乘 gamma,加 beta)

const __half* beta = ...; // 偏置 [K]
DirectLoad<__half, float> load(x, K);
AffineStore<float, __half, true, true> store(y, K, gamma, beta);
// ^^^^ ^^^^
// do_scale=true do_center=true
// kernel 中的 load/store 调用完全相同!
load.load<8>(buf, row, col);
// ... 归一化计算 ...
store.store<8>(buf, row, col); // 存储: float→half, 乘 gamma, 加 beta

纯归一化(无 affine 参数)

DirectLoad<float, float> load(x, K);
DirectStore<float, float> store(y, K); // 或 AffineStore<float,float,false,false>
// ^^^^^ ^^^^^
// SRC=float DST=float (无类型转换)
// kernel 中的调用依然相同
load.load<4>(buf, row, col);
// ... 归一化计算 ...
store.store<4>(buf, row, col);

三种场景下,kernel 代码完全不变。 只有 Load/Store 对象的创建方式不同。这就是 io.cuh 分离”数据搬运”和”计算逻辑”的设计价值。

Norm Kernel模块#

有了IO和Reduce,我们就可以很好的实现Norm Kernel。

先看两个 kernel 的伪代码对比:

RMSNorm: LayerNorm:
for each row: for each row:
sq_sum = 0 sum = 0; sq_sum = 0
for each col: for each col:
v = load(x[row][col]) v = load(x[row][col])
sq_sum += v * v sum += v
inv_rms = rsqrt(sq_sum/K + eps) sq_sum += v * v
for each col: mean = sum / K
out = v * inv_rms * gamma[col] var = sq_sum/K - mean²
store(y, out) inv_std = rsqrt(var + eps)
for each col:
out = (v - mean) * inv_std
* gamma[col] + beta[col]
store(y, out)
二者绝大部分代码都相同

行遍历、列加载、累加、reduce、归一化、这些骨架完全一致

因此我们完全可以拆出一个Kernel来把这些骨架放进去,然后用Stats结构体去描述我们需要的参数量。

Dispatch#

template <typename T, int pack_size>
static void rms_norm_dispatch(..., int N, int K, float eps) {
if (K <= 1024)
RMSNormWarpDispatch<T, pack_size>::launch(...); // 策略一
else
launch_rms_smem<T, pack_size, 256>(...); // 策略二
}

这里根据归一化维度 K 的大小,选择完全不同的 kernel 策略。

K 范围策略核心思想
K ≤ 1024WarpImpl一个 warp group处理一行,一个 block 同时处理 4 行
K > 1024BlockSMemImpl整个 block(256 线程)协作处理一行,shared memory 缓存避免重复读
┌──────────────┐
│ reduce.cuh │ ← warp/block 规约原语
└──────┬───────┘
│ 被 Stats 调用
┌────────────┼────────────┐
│ ▼ │
┌─────────┴─────┐ ┌──────────┐ ┌───┴──────────┐
│ io.cuh │ │ Stats │ │ io.cuh │
│ DirectLoad │ │ (RMS/LN) │ │ AffineStore │
└───────┬───────┘ └────┬─────┘ └─────┬─────────┘
│ │ │
└──────────────┼─────────────┘
┌────────────────────────┐
│ norm_kernel.cuh │
│ NormWarpImpl / │
│ NormBlockSMemImpl │
└────────────────────────┘

WarpImpl#

使用场景是当K比较小的时候(hidd_size比较小),N可以任意。

核心思路

既然 K 小到可以塞进一个 warp 的寄存器,就不需要整个 block 协作。 把 block 按 warp 拆分,每个 warp 处理独立的一行。

blockDim = (32, 4) ← 32 线程的 warp group,4 个 group 叠在一起
gridDim = (N/4, 1)
block (32, 4)
├── warp group 0 (lane 0..31, threadIdx.y=0) → row 0
├── warp group 1 (lane 0..31, threadIdx.y=1) → row 1
├── warp group 2 (lane 0..31, threadIdx.y=2) → row 2
└── warp group 3 (lane 0..31, threadIdx.y=3) → row 3

BlockSMemImpl#

适用场景是当K比较大的时候(大hidden_size,例如4096, 8192)

朴素实现需要读取输入X两次,第一次做统计量,第二次做归一化。

优化方案也很简单,第一次统计之后放在Shared Memory中,第二次归一化的时候直接从共享内存取就可以。

第一趟 (Global → SMem + Accumulate):
Global Mem ──read──→ [Shared Memory] ──→ [Thread Accumulator]
col 0..K-1 (sum / sq_sum)
Block Reduce → stat:
[各线程的 acc] ──reduce──→ s_stat (只有 warp0 持有)
Broadcast (通过 shared memory):
if (threadIdx.x == 0) s_stat = stat;
__syncthreads(); ← 第一次同步:等 warp0 写完
stat = s_stat; ← 全部线程读到
第二趟 (SMem → Normalize → Global):
[Shared Memory] ──read──→ normalize ──store──→ Global Mem y

RMS和LN模块#

rms_norm.cuhlayer_norm.cuh 的职责是定义Stats。

norm_kernel 需要什么? rms_norm.cuh 提供什么?
───────────────────────── ───────────────────────
Stats::accum_t → 累加器类型 float (只存平方和)
Stats::stat_t → 统计量类型 float (inv_rms)
Stats::init(acc) → 初始化为 0 置零
Stats::accumulate → 每读一个元素 a += v²
Stats::warp_reduce → warp 内合并 warp_reduce_sum
Stats::block_reduce → block 内合并 block_reduce_sum
Stats::compute → 算出最终统计量 rsqrt(a/K + eps)
Stats::normalize → 就地归一化 v * inv_rms
加上 Dispatch 层:
K ≤ 1024 → NormWarpImpl<..., RMSNormStats, ...>
K > 1024 → NormBlockSMemImpl<..., RMSNormStats, ...>
再封装 Host API:
rms_norm_forward(x, y, gamma, N, K, eps, is_fp16)

对于Stats来说:

struct Stats {
// 模式一:void — 原地修改,状态累积
static void init(accum_t& a); // 引用参数
static void accumulate(accum_t& a, const T* vals, int n); // 引用参数
static void normalize(T* vals, stat_t s, int n); // 指针参数
// 模式二:return value — 纯函数,产生新值
static accum_t warp_reduce(accum_t a); // 值参数 + 返回值
static accum_t block_reduce(accum_t a); // 值参数 + 返回值
static stat_t compute(accum_t a, int K, float eps); // 值参数 + 返回值
};

具体完整的代码可以查看仓库,至此所有关于Norm系列的内容讲解完毕。

支持与分享

如果这篇文章对你有帮助,欢迎分享给更多人或赞助支持!

赞助
LeetGPU习题07:Norm系列代码实现
https://dlog.com.cn/posts/leetgpu07/norm/
作者
杜子源
发布于
2026-06-12
许可协议
CC BY-NC-SA 4.0
Profile Image of the Author
杜子源
都是风景,幸会
公告
请狠狠地打赏我,打赏一次,爆更一篇!!
音乐
封面

音乐

暂未播放

0:00 0:00
暂无歌词
分类
标签
站点统计
文章
29
分类
8
标签
11
总字数
93,152
运行时长
0
最后活动
0 天前

目录