把归一化塞进 GEMM 与 Attention:让 LayerNorm 和 RMSNorm 接近“免费”

2026-07-10 38 预计阅读时间: 1 分钟
来源: pytorch.org AI 摘要 Original link

Disclaimer: This article is an AI-assisted summary. Read it together with the original source when precision matters. The summary may omit context, version differences, or edge cases and is not official documentation.

预计阅读时间:10 分钟

LayerNorm 和 RMSNorm 的浮点运算量并不算大,却经常成为模型推理与训练链路中的明显开销。问题通常不在算术,而在数据搬运:一个独立归一化算子需要读取输入、计算统计量、写回结果,后续 GEMM 或 Attention 又要重新读取这些结果。

本文讨论的优化方向,是把常见归一化操作融合进 GEMM 和 Attention 内核,让中间张量尽量停留在寄存器或片上存储中。来源工作给出了多种新的内核融合技术,并报告了显著加速;其核心价值不只是减少一次 kernel launch,而是压缩整条执行链上的全局内存流量。

为什么一个“小算子”会拖慢整条链路

以 RMSNorm 后接线性投影为例,未融合的执行过程可以抽象成:

X --读取--> RMSNorm --写回--> X_norm --再次读取--> GEMM(X_norm, W) --> Y

RMSNorm 对每一行执行:

rms = sqrt(mean(x²) + eps)
y   = x / rms * weight

这里包含平方和归约、缩放以及可选的逐元素权重。相对于大型矩阵乘法,计算量很小,但独立实现会产生几个现实成本:

  • 启动额外 GPU kernel;
  • 从全局内存读取输入并写出归一化结果;
  • GEMM 随后再次读取中间结果;
  • 在行宽较大时,归约可能跨越多个线程块,带来同步和调度开销;
  • 小批量或短序列场景下,kernel launch 和访存成本更难被大规模计算掩盖。

因此,“免费归一化”不是指归一化不再做计算,而是让它的计算和访存成本被相邻的高吞吐内核吸收。

融合真正省掉了什么

将归一化融合进 GEMM,可以把流程改写为:

X --读取一次--> 行统计量 + 缩放 --> 直接送入 GEMM --> Y

如果数据布局和 tile 划分允许,内核可以在加载 GEMM 输入时同时累计平方和或均值、方差,然后完成缩放,再把结果交给矩阵乘法主循环。这样做通常能减少中间张量的落盘以及独立 kernel 的调度开销。

Attention 里的机会更大。归一化、QKV 投影、注意力计算和输出投影之间往往存在连续的数据依赖。把这些阶段组织成更大的融合内核,也就是常说的 megakernel,可以让部分中间状态留在 GPU 片上。不过,内核越大,工程约束也越严格:

  • 寄存器使用量上升可能降低 occupancy;
  • 不同序列长度、隐藏维度和数据类型需要不同 tile;
  • 跨线程块归约需要额外的协作机制;
  • 动态 shape、mask、dropout 和训练反向传播会扩大实现范围;
  • 某些规模下,拆分内核反而更容易取得更高并行度。

来源代码目录中体现了两个值得关注的方向:多 CTA 协作的归一化融合,以及面向 Attention 的大型融合内核。它们说明归一化融合并不限于“一个线程块处理一行”的简单情况,还需要处理行宽过大或融合链路较长时的并行协作问题。

可以这样实践:先用 PyTorch 验证融合机会

下面的示例不复现来源中的定制 GPU 内核,而是一个可直接运行的诊断程序:比较 eager 模式与 torch.compile 对“RMSNorm + Linear”链路的处理效果。它适合用来判断当前模型、shape 和硬件上是否存在值得继续开发定制内核的空间。

运行前需要安装支持 CUDA 的 PyTorch 2.x。将 MKN 改成业务中的 token 数、隐藏维度和投影维度。

import time
import torch
import torch.nn as nn

assert torch.cuda.is_available(), "This benchmark requires a CUDA GPU"
torch.manual_seed(0)
torch.set_float32_matmul_precision("high")

M, K, N = 4096, 4096, 4096
DTYPE = torch.bfloat16
DEVICE = "cuda"

class RMSNormLinear(nn.Module):
    def __init__(self, hidden_size: int, output_size: int, eps: float = 1e-6):
        super().__init__()
        self.weight = nn.Parameter(torch.ones(hidden_size, device=DEVICE, dtype=DTYPE))
        self.linear = nn.Linear(
            hidden_size, output_size, bias=False, device=DEVICE, dtype=DTYPE
        )
        self.eps = eps

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        variance = x.float().square().mean(dim=-1, keepdim=True)
        normalized = x * torch.rsqrt(variance + self.eps).to(x.dtype)
        return self.linear(normalized * self.weight)

model = RMSNormLinear(K, N).eval()
x = torch.randn(M, K, device=DEVICE, dtype=DTYPE)
compiled_model = torch.compile(model, mode="max-autotune", fullgraph=True)

@torch.inference_mode()
def benchmark(fn, warmup=20, iterations=100):
    for _ in range(warmup):
        fn(x)
    torch.cuda.synchronize()

    start = time.perf_counter()
    for _ in range(iterations):
        fn(x)
    torch.cuda.synchronize()
    return (time.perf_counter() - start) * 1000 / iterations

with torch.inference_mode():
    eager_output = model(x)
    compiled_output = compiled_model(x)

torch.testing.assert_close(compiled_output, eager_output, rtol=2e-2, atol=2e-2)
print(f"eager:    {benchmark(model):.3f} ms")
print(f"compiled: {benchmark(compiled_model):.3f} ms")

这个实验不能证明编译器一定生成了“归一化直接进入 GEMM”的单一内核。应配合 profiler 检查 kernel 数量、执行时间和显存流量:

nsys profile \
  --trace=cuda,nvtx,osrt \
  --sample=none \
  --output=rmsnorm_linear \
  python benchmark.py

nsys stats rmsnorm_linear.nsys-rep

观察重点包括:编译前后 kernel launch 数是否下降、是否仍然物化完整的归一化输出、归一化 kernel 的耗时是否消失或缩短,以及 GEMM 性能是否因融合后的资源压力而退化。

定制内核不能只比较最终延迟

如果决定采用专用融合内核,基准测试至少要覆盖实际生产 shape,而不是只测一个方阵。建议建立如下测试矩阵:

维度 建议覆盖范围
Token 数 单 token、小批量、长序列
Hidden size 模型实际使用的全部维度
数据类型 FP16、BF16,以及需要时的 FP8 路径
归一化 LayerNorm、RMSNorm、是否带 bias 或残差
Attention causal/non-causal、不同 mask 和序列长度
执行模式 推理、训练前向、反向传播

还要单独验证数值误差。归约次序改变后,LayerNorm 的均值和方差、RMSNorm 的平方和都可能出现浮点差异。低精度输入通常应使用更高精度累加,并针对极小方差、大幅值输入和非连续张量设置测试。

性能分析则至少记录:

  • 端到端延迟,而不只是单 kernel 时间;
  • DRAM 读写字节数;
  • kernel launch 数量;
  • 寄存器和共享内存占用;
  • occupancy 与实际吞吐;
  • 编译时间、二进制体积和 shape 覆盖率。

落地时的判断清单

归一化融合最适合结构稳定、shape 集中、调用频繁的模型路径。若一个 RMSNorm 总是紧邻 QKV 投影,或者 LayerNorm 总是进入固定尺寸的 GEMM,那么专用内核更容易摊薄开发和维护成本。

采用前可以依次确认:

  1. Profiler 是否证明该归一化链路受内存带宽或 kernel launch 限制。
  2. 中间归一化结果是否只有一个直接消费者,能否避免写回全局内存。
  3. 生产 shape 是否足够集中,值得维护专门的 tile 和调度策略。
  4. 融合后寄存器压力是否导致 GEMM 或 Attention 主体性能下降。
  5. 动态 shape、mask、训练反向和数值容差是否有明确的回退路径。
  6. 未命中优化 shape 时,是否能可靠退回框架原生实现。

这类优化的目标不是追求最大的单个 kernel,而是让数据在正确的位置停留更久。只有当减少的内存往返和调度成本超过融合引入的资源压力与工程复杂度时,LayerNorm 和 RMSNorm 才真正接近“免费”。


相关推荐