在Transformer类模型的推理性能分析中,经常会出现一个反直觉的现象:模型里计算量最大的矩阵乘法并不是最耗时的部分,反而是LayerNorm、GeLU、Dropout这类看起来计算量很小的算子,占用了相当大的时间比例。原因很简单,这类算子是典型的访存密集型操作,它们的瓶颈不在计算而在数据搬运。算子融合(Operator Fusion)就是针对这个问题最直接有效的优化手段,把LayerNorm和GeLU这类相邻的小算子合并成一个内核一次执行,通常能带来明显的延迟下降。本文以LayerNorm与GeLU的融合为例,把原理和实现讲透。

为什么逐算子执行会慢:理解算子融合的动机
要理解算子融合的价值,先要看清楚不融合时的开销在哪里。深度学习框架在执行一个网络时,是以算子为单位调度GPU内核的。每执行一个算子,都需要经历内核启动、从显存读取输入、计算、把结果写回显存这四个阶段。以LayerNorm后接GeLU的组合为例,标准流程是:先启动一个LayerNorm内核,从显存读入x,计算均值方差和归一化结果,写到中间张量y;再启动一个GeLU内核,从显存读入y,做激活计算,写到最终输出z。
这里的问题在于,y这个中间结果完全没必要落地到显存。对于混合精度模型,x通常是fp16或bf16类型,隐藏维度假设为4096,一个token在LayerNorm和GeLU之间的中间数据就要读写各8KB。当batch和序列长度放大后,这部分访存量会成倍增长。而LayerNorm本身每行只需要做一次均值和方差的归约,再加上几个逐元素运算,计算量非常小,GPU的计算单元大部分时间在等数据从显存搬过来。内核启动本身也有开销,一次内核launch大约在几微秒量级,当算子数量很多、每个算子执行时间又很短时,launch开销的占比会非常可观。
算子融合的核心思路是:把LayerNorm和GeLU合并成一个内核,在处理每一行数据时,先在寄存器或共享内存中完成均值方差计算,得到归一化结果后不写回显存,而是直接在寄存器里套用GeLU公式,最后只把最终结果写回一次。这样中间结果的两次显存读写(一次写、一次读)就完全消失了,内核launch次数也从两次降为一次。对于访存密集型算子,这种优化几乎能把执行时间压缩到原来的一半以下。
PyTorch中的手写融合实现与验证
在PyTorch的原生实现中,torch.nn.LayerNorm和torch.nn.functional.gelu是两个独立的操作。我们先看未融合的写法,并用性能测试确认瓶颈:
import torch
import torch.nn.functional as F
x = torch.randn(8, 4096, 4096, dtype=torch.float16, device="cuda")
ln = torch.nn.LayerNorm(4096).cuda().half()
# 未融合:两个独立的CUDA内核
def baseline(x):
return F.gelu(ln(x))
# 用CUDA Event精确计时
starter = torch.cuda.Event(enable_timing=True)
ender = torch.cuda.Event(enable_timing=True)
for _ in range(10): # 预热
baseline(x)
torch.cuda.synchronize()
starter.record()
for _ in range(100):
baseline(x)
ender.record()
torch.cuda.synchronize()
print("baseline:", starter.elapsed_time(ender) / 100, "ms")
这段代码在A100级别的卡上跑,baseline每次大约0.5毫秒左右(具体数值随硬件变化)。接下来用PyTorch的TorchScript trace机制做一个最简单的融合,框架在编译时会自动把可以合并的逐元素操作合并到一起:
fused_fn = torch.jit.trace(lambda t: F.gelu(ln(t)), x)
starter.record()
for _ in range(100):
fused_fn(x)
ender.record()
torch.cuda.synchronize()
print("fused:", starter.elapsed_time(ender) / 100, "ms")
trace后的版本通常会有一定提升,但TorchScript的融合能力有限,对于包含归约操作(LayerNorm需要算均值和方差)的算子,它经常无法完成理想的融合。这时候更可靠的做法是手写融合内核,PyTorch提供了torch.compile,在2.x版本后其底层由Triton驱动,融合能力强大得多:
@torch.compile
def fused_ln_gelu(x, weight, bias, eps=1e-5):
y = F.layer_norm(x, x.shape[-1:], weight, bias, eps)
return F.gelu(y)
out = fused_ln_gelu(x, ln.weight, ln.bias)
无论采用哪种方式,上线前必须做数值验证。融合前后用torch.allclose对比结果,注意要设置合理的容差,因为融合内核内部的累加顺序可能与原始实现不同,浮点结果会有微小差异,这在数值上是正常的:
ref = baseline(x) out = fused_ln_gelu(x, ln.weight, ln.bias) print(torch.allclose(ref, out, atol=1e-3, rtol=1e-3)) print((ref.float() - out.float()).abs().max().item()) # 查看最大绝对误差
用Triton手写融合内核:结构与关键细节
如果想进一步压榨性能,或者需要部署在不依赖TorchInductor的环境中,可以直接用Triton写融合内核。Triton是OpenAI开源的GPU编程语言,写法接近NumPy,同时能利用shared memory和warp级别的并行。下面是一个完整的LayerNorm+GeLU融合实现:
import triton
import triton.language as tl
@triton.jit
def ln_gelu_kernel(
X, W, B, Y,
stride, N, eps,
BLOCK: tl.constexpr,
):
row = tl.program_id(0)
# 计算当前行的起始地址
x_ptr = X + row * stride
y_ptr = Y + row * stride
offs = tl.arange(0, BLOCK)
mask = offs < N
x = tl.load(x_ptr + offs, mask=mask, other=0.0).to(tl.float32)
# LayerNorm:先算均值
mean = tl.sum(x, axis=0) / N
diff = tl.where(mask, x - mean, 0.0)
var = tl.sum(diff * diff, axis=0) / N
rstd = 1.0 / tl.sqrt(var + eps)
w = tl.load(W + offs, mask=mask).to(tl.float32)
b = tl.load(B + offs, mask=mask).to(tl.float32)
y = diff * rstd * w + b
# GeLU:直接作用在寄存器中的结果上,无显存读写
# tanh近似版本:0.5 * y * (1 + tanh(sqrt(2/pi) * (y + 0.044715 * y^3)))
inner = 0.7978845608 * (y + 0.044715 * y * y * y)
gelu_out = 0.5 * y * (1.0 + tl.math.tanh(inner))
tl.store(y_ptr + offs, gelu_out.to(tl.float16), mask=mask)
注意几个关键细节。第一,计算均值和方差时一定要把数据转成fp32再累加,输入是fp16时直接在fp16域做归约会损失精度,尤其是方差计算涉及平方操作,误差会被放大。第二,GeLU有两种实现:精确的erf版本和tanh近似版本,tanh版本在GPU上通常更快,与PyTorch默认实现存在约1e-3量级的差异,如果下游对精度敏感,需要改用tl.math.erf。第三,BLOCK大小要覆盖整行,Triton要求tl.arange的长度是2的幂,所以当N不是2的幂时必须用mask处理越界,加载时用other=0.0填充,但归约计算前要记得把填充值从diff中去掉,否则方差会算错。
调用侧的wrapper代码如下,处理好grid的划分和BLOCK的选择:
def ln_gelu(x, weight, bias, eps=1e-5):
x = x.contiguous()
y = torch.empty_like(x)
M, N = x.numel() // x.shape[-1], x.shape[-1]
BLOCK = triton.next_power_of_2(N)
ln_gelu_kernel[(M,)](
x, weight, bias, y,
x.stride(0), N, eps,
BLOCK=BLOCK,
)
return y
这个实现在隐藏维度4096、输入尺寸较大时,相比未融合版本通常能取得2倍左右的加速,且随batch增大提升更明显,因为访存节省的绝对量与数据量成正比,而launch开销占比则随数据量增大而相对下降。
融合实践中的常见坑与选型建议
第一个常见的坑是盲目融合。并非所有算子组合都适合融合,融合的前提是数据流在融合区域内是局部的:LayerNorm按最后一维归约,GeLU逐元素作用,两者天然可以按行融合。但如果中间夹着一个需要全局信息的操作,比如softmax或跨token的attention,强行融合反而会增加实现复杂度却拿不到访存收益。判断标准很简单:融合后的内核中,每个线程块能否独立处理完自己负责的数据而不需要跨块通信。
第二个坑是忽略自动融合框架的能力。TensorRT-LLM、vLLM、torch.compile这些框架在算子融合上已经做了大量工程化工作。以torch.compile为例,默认的mode就会触发算子融合,还可以通过torch._inductor.config调整融合策略;vLLM则直接内置了RMSNorm加SwiGLU等Transformer典型组合的融合内核。如果你的场景是标准Transformer推理,优先用这些成熟方案,自己手写Triton内核只在框架无法覆盖的定制结构中才有必要。
第三个坑是验证方式不当。融合内核改动了浮点运算的顺序和精度策略,逐bit对比必然失败,正确的验证方法是统计相对误差和下游任务指标。同时要做好fallback设计,在异常输入维度(比如隐藏维度超出预设BLOCK上限)时回退到原生实现,保证鲁棒性。另外建议把Triton内核做成可缓存的形式,首次编译耗时不小,如果每次进程启动都重新编译,会拖慢服务冷启动速度。
总结一下,LayerNorm与GeLU融合之所以能提速,本质是消除了中间张量的显存往返和多余的内核launch,把访存密集的执行模式改造成接近一次读一次写的流式处理。在实际项目中,建议先用profiler确认瓶颈确实在这类小算子上,再优先尝试torch.compile或推理引擎的内置融合,最后才考虑手写Triton内核,这样能把工程投入控制在收益最大的范围内。