LeetGPU习题07:Norm系列代码实现
Pytorch调用
PyTorch的调用接口是高度统一的,它们都继承自nn.Module,核心都是沿着某个维度计算统计量,之后标准化,可选仿射变换。
区别在于沿着哪个轴做归一化,以及计算哪些统计量
理解了这个本质,剩下的就是查找参数,看维度。
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_normLayerNorm
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 torchimport 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一致。
| 特性 | BatchNorm | LayerNorm | RMSNorm |
|---|---|---|---|
| 归一化轴 | 沿 batch 和空间维度 (每个通道独立) | 沿特征维度 (每个样本独立) | 沿特征维度 (每个样本独立) |
| 依赖 batch | 是 | 否 | 否 |
| 可学习参数 | weight + bias (通道数) | weight + bias (归一化形状) | 仅 weight |
| 适用场景 | CNN、CV | NLP、Transformer | 大模型 |
| PyTorch 接口 | BatchNorm1d/2d/3d | LayerNorm | RMSNorm |
Triton实现
在了解了Pytorch这一层级的实现之后,我们要手动在Triton中实现LN和RMS。
下边我们核心先讲解LN和RMS的Triton版本,BN我就先忽略,大家可以自行学习。
归一化的视角
无论是LN还是RMS,输入都是一个二维矩阵,其中N是所有token的总行数,K是特征维度。
每一行的计算是完全独立的。 对于第n行,算子聚合该行内所有K个元素,算出指定的统计量,再用统计量去逐元素变换该行。
这是经典的Reduce+Broadcast模式
具体来说,LayerNorm需要两步聚合,分别计算均值和方差;RMSNorm只进行一次平方和聚类。
在GPU上,Reduce是代价最高的部分,Broadcast是element wise操作。很明显,Norm操作是带宽密集型操作。
LayerNorm
我在这里参考了flash-attenion官方的实现: flash-attention
实际上,除了前向传播,在写绝大多数算子的时候,也要考虑反向传播,后续我会写一个完整的macroTorch,来完整的学习是如何进行前向传播与反向传播的。
关于自动调优,可以写一个函数来囊括进去常用的内容。 同时我在这里保存了mean和rstd,这样能够给未来的反向传播直接复用,避免重新进行规约计算,因此可以看到我显式的进行了两次保存。
并且还有一个设计上的巧思,即步长参数化,而不是假设内存严格连续,这可以让kernel实现更加复杂的张量布局。
import torchimport torch.nn as nnimport tritonimport triton.language as tlfrom 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.jitdef 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 = Nonetry: 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_reportexcept 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")性能对比如下:

RMSNorm
LayerNorm需要计算两个统计量:mean和std。 mean就得涉及到规约,而RMSNorm只计算一个:
rms = sqrt(E[x^2] + eps)y = x / rms * γ核心代码如下:
@triton.autotune( configs=autotune_configs(), key=["N"],)@triton.jitdef 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在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.46xFP32下加速比在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左右,一个的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. 规约与归一化分离
朴素实现中,规约和归一化交织在一起,例如需要先完整遍历一次数据求出 ,再遍历第二次求方差,第三次才做归一化。这不仅 I/O 开销大,而且求方差时容易产生严重的数值精度问题。
朴素求方差的公式是:
这要求同时累积 和 。当数据的均值较大而方差本身很小时, 和 是两个非常接近的大数,相减会导致大量有效数字抵消,这叫做 catastrophic cancellation,结果误差极大。
Welford 算法提供了一种数值稳定的在线计算方式,它不需要保存 ,而是维护一个与均值无关的修正平方和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 算法天然支持单趟遍历,一边读取数据一边更新 mean 和 m2。
GPU 上求全局均值/方差需要跨线程块规约。Welford 具备可合并性 这允许我们先让每个 warp / block 独立做局部 Welford,然后只归并这几个状态.
一旦得到全局 和 ,归一化就变成了一个完全无依赖的逐元素操作:(x - μ) * inv_std。
这可以很方便地和其他逐元素操作(如 、残差加法、激活函数)融合到一个 kernel 里,形成端到端的“大融合算子”。
分离的设计让规约成为一次全局通信+少量计算,归一化成为纯本地并行计算,边界清晰,利于编译器和手写 kernel 深度优化。
3. 策略与算法分离
| 策略 | 适用条件 | 线程组织 | 核心思路 |
|---|---|---|---|
| WarpImpl | K ≤ 1024 | 2D block | 一个 warp group 一行,block 内多行并行 |
| BlockSMemImpl | K > 1024 | 1D block + dynamic shared memory | 数据缓存到 shared memory,省去第二次 global 读 |
| BlockUncachedImpl | SMemImpl 放不下 | 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.cuhReduce模块
提供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 ... lane31mask=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.128template <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 指令。
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]); // 自动类型转换 }};SRC 和 DST 两个模板参数是这一层的核心设计:
SRC= 显存中的存储类型(__half用于 FP16 模型)DST= 寄存器中的计算类型(float用于 FP32 精度计算)
GPU 显存: fp16 ──DirectLoad<__half, float>──→ 寄存器: fp32 自动调用 __half2float当然是,容易忘!
DirectStore 是 DirectLoad 的镜像,方向相反,逻辑完全对称:
pack = *reinterpret_cast<const Pack<SRC, N>*>(src + offset * N); // Load: Global → Pack*reinterpret_cast<Pack<DST, N>*>(dst + offset * N) = pack; // Store: Pack → GlobalAffineStore编译期分支消除
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 和 StoreDirectLoad<__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, 乘 gammaLayerNorm (乘 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 ≤ 1024 | WarpImpl | 一个 warp group处理一行,一个 block 同时处理 4 行 |
| K > 1024 | BlockSMemImpl | 整个 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 3BlockSMemImpl
适用场景是当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 yRMS和LN模块
rms_norm.cuh 和 layer_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_sumStats::block_reduce → block 内合并 block_reduce_sumStats::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系列的内容讲解完毕。
支持与分享
如果这篇文章对你有帮助,欢迎分享给更多人或赞助支持!