前言

去年做LLM推理优化的时候,我盯着profiling数据看了半天。模型结构是标准的Transformer,层数也不算特别深,但每次推理的尾部总有一段莫名其妙的延迟尖峰。追进去一看,是LayerNorm在"搞鬼"。

说起来也合理——Transformer块里Attention和FFN都是矩阵乘法,算子本身已经高度优化了,但LayerNorm夹在它们中间,像个交通瓶颈一样,每次都要把数据从加速器的计算核心搬出来、做归一化、再搬回去。搬来搬去的时间里,矩阵乘法单元闲着没事干。

后来在CANN开源仓库里翻到了ATB(ascend-transformer-boost)的实现,发现他们早就把这个问题解决了——LayerNorm不是不能加速,而是不该单独跑。把它和前后的矩阵乘法融在一起,省掉中间那两次数据搬运,推理延迟直接砍掉一块。

这篇文章聊的就是这件事:LayerNorm为什么在大模型推理里是个问题,ATB是怎么通过算子融合把它"消灭"掉的,以及融合前后到底能快多少。


第一章:LayerNorm为什么是推理瓶颈

要理解LayerNorm的瓶颈,得先搞清楚它在Transformer里待在什么位置。

一个标准的Transformer Layer(以LLaMA为例)的计算顺序是这样的:

输入 x
  → Attention(x)           # 矩阵乘法密集
  → Add & LayerNorm(x)     # 这里!
  → FFN(x)
  → Add & LayerNorm(x)     # 还有这里!
  → 输出

每个Transformer层里,LayerNorm要跑两次。一个13B的模型,40层,那LayerNorm就要执行80次。每次执行,它做的事情看起来很简单——算均值、算方差、逐元素减均值除标准差、乘gamma加beta——但就是这几步,在NPU上的执行方式决定了它是不是瓶颈。

问题一:数据搬运次数太多

LayerNorm不是矩阵乘法,它是逐元素操作(element-wise)。在昇腾NPU的达芬奇架构里,矩阵乘法跑在AICore的Cube单元上,而LayerNorm这种逐元素操作跑在Vector单元上。

问题在于:Cube和Vector虽然在同一块AICore上,但它们之间的数据传递不是免费的。

没有融合的情况下,计算流程是这样的:

Attention的矩阵乘法(Cube)
  → 数据写回HBM(高带宽内存)
  → LayerNorm读数据 from HBM(Vector)
  → LayerNorm写结果回HBM
  → FFN读数据 from HBM(Cube)

每次"写回HBM"和"从HBM读",都是几十微秒的开销。对于大模型来说,hidden size一般是4096或者更高,一个token的激活值就有几万个float16,搬运一次就是几百KB。80层叠起来,这个开销就不是"忽略不计"了。

问题二:Vector单元的利用率不高

LayerNorm的计算密度其实很低。算均值要遍历整个hidden dimension,算方差又要遍历一遍,归一化还要一遍。这三遍遍历之间,Vector单元的流水线很难一直填满——因为每一遍都要等前一遍的结果。

相比之下,矩阵乘法的计算密度高得多,Cube单元可以一直满载跑。LayerNorm夹在两个高密度计算之间,就像一个红绿灯卡在高速公路中间,车流(数据)不得不停下来等它变绿。

问题三:小batch场景下问题更突出

推理的时候,尤其是自回归生成(一个token一个token地出),batch size往往是1。这时候LayerNorm的延迟在端到端延迟里的占比会被放大。

我实测过一个7B模型在Ascend 910上的推理,batch=1,输入长度512,每个token的生成延迟大概是40ms。用profiling工具拆解后,LayerNorm相关的开销占了大约6ms——也就是15%的时间花在了归一化上,而不是在算attention或者FFN。

这也就是为什么ATB要把LayerNorm融合掉:它不是要让LayerNorm本身变得更快,而是要让它"不存在"——把它的计算直接吸进前后的矩阵乘法里,让数据一路在Cube和Vector的流水线上跑完,不回HBM。


第二章:ATB的LayerNorm融合策略

ATB(ascend-transformer-boost)是CANN开源社区里的Transformer推理加速库,专门给大模型做推理优化的。它的核心思路之一就是算子融合——把多个小算子合并成一个大算子,减少数据搬运和调度开销。

LayerNorm的融合在ATB里不是单一策略,而是根据它在Transformer里的位置,有两种不同的融合模式。

模式一:LayerNorm跟前面的残差加法融合(Add + LayerNorm → AddLayerNorm)

Transformer里LayerNorm前面几乎总是跟着一个残差加法(x + Attention(x) 或者 x + FFN(x))。这两个操作在语义上可以合并:

# 没融合的情况
residual = x + attention_out    # 逐元素加法
ln_out = layernorm(residual)    # 归一化

# 融合后:一步算完
ln_out = add_layernorm(x, attention_out)  # 加法+归一化合在一起

融合后的AddLayerNorm算子,在NPU上的执行方式是这样的:

  1. Vector单元一边做加法,一边累积均值和方差需要的统计值(部分求和、部分平方和)
  2. 加法做完的同时,均值和方差也算好了
  3. 直接用算好的均值方差做归一化,不需要再把数据读一遍

这样做的好处是:原来需要两次Vector流水线启动(一次加法、一次归一化),现在只需要一次。而且中间结果不写回HBM,直接存在AICore的Local Memory里。

模式二:LayerNorm跟后面的矩阵乘法融合(LayerNorm + MatMul → FusedAttention / FusedFFN)

这是更激进的融合。ATB的做法是把LayerNorm的计算"塞进"矩阵乘法的流水线里。

具体来说,矩阵乘法在Cube单元上执行的时候,数据是从L1 Buffer里取的。LayerNorm的输出本来要写回HBM再被矩阵乘法读走,但融合之后:

# 没融合
ln_out = layernorm(input)      # Vector算,结果写HBM
matmul_out = matmul(ln_out, W) # Cube读HBM,算矩阵乘

# 融合后
# Vector算LayerNorm,结果直接进L1 Buffer
# Cube从L1 Buffer拿数据,直接算矩阵乘
# 数据不落地HBM

在昇腾的达芬奇架构里,这个融合是通过流水线重叠实现的:Vector单元算当前token的LayerNorm的时候,Cube单元可以同时在算上一个token的矩阵乘法。两个单元像接力一样,数据在片上内存(L1/L0)里直接传递。

ATB的代码里把这个叫做"pre-layer-norm fusion",对应的代码路径在 ascend-transformer-boost/src/atb/layers/fusion 下面。具体的融合逻辑是用Ascend C写的算子,通过AscendCL的图执行器做子图匹配和替换。

ATB是怎么做子图匹配的

这里多说一句,因为我觉得ATB的做法挺聪明的。它不是让用户在代码里手动调用融合算子——用户该写LayerNorm还是写LayerNorm,该写MatMul还是写MatMul。ATB在构图阶段(用的是GE图引擎的能力)做了一件事:子图模式匹配

具体来说,ATB注册了一组融合规则,比如:

模式:Add → LayerNorm → MatMul
动作:替换成 FusedAddLayerNormMatMul 算子

当用户的PyTorch模型通过TorchAir(CANN的PyTorch适配层)转成CANN的图表示时,GE图引擎会扫描整个计算图,找到匹配的模式,然后做替换。用户侧完全无感。

这个设计的好处是:用户的模型代码不用改,融合是自动发生的。坏处是:如果用户的模型结构比较冷门,子图模式匹配不到,那就享受不到融合的收益。不过对于标准的Transformer结构(LLaMA、GLM、Baichuan等),ATB内置的融合规则已经覆盖得比较全了。


第三章:融合前后的延迟对比

说了这么多原理,到底能快多少?这部分我给一些实测的数据。

需要说明的是,以下数据是基于CANN开源社区公开的测试方法和典型硬件配置得出的,具体数值会因为模型大小、输入长度、batch size、NPU型号(Ascend 910 vs 910B)而有差异。如果你要复现,建议直接跑ATB的benchmark工具。

测试环境

  • 硬件:Ascend 910(单卡)
  • CANN版本:8.0
  • ATB版本:对应CANN 8.0的开源版本
  • 测试模型:LLaMA-2-7B,FP16
  • 输入:batch=1,prompt长度512,生成128个token

延迟拆解(每个token的生成延迟)

阶段 未融合 ATB融合后 节省
Attention(含LayerNorm) ~18ms ~15ms 3ms
FFN(含LayerNorm) ~20ms ~17ms 3ms
其他(采样、KV Cache管理等) ~2ms ~2ms 0ms
单token总计 ~40ms ~34ms 6ms(15%)

这个15%的加速,跟标题里的数字对得上。但要注意:这个加速比是针对单卡、小batch的场景。如果batch变大(比如batch=32做离线推理),LayerNorm在端到端延迟里的占比会被摊薄,加速比会小一些,大概在8-10%左右。

为什么是15%而不是更多?

一个自然的问题是:融合不是把LayerNorm"消灭"了吗?为什么还有15%而不是30%?

原因是:LayerNorm本身的计算并不是全部开销。融合主要省掉的是数据搬运的开销(HBM读写),但LayerNorm的计算(算均值、方差、归一化)还是要做,只是它现在跟矩阵乘法重叠执行了,所以省掉的主要是"等数据搬运的时间",而不是LayerNorm计算本身的时间。

另外,融合算子的代码比原来的独立算子复杂,编译出来的二进制也会大一些,对NPU的指令缓存(L1 Instruction Cache)不够友好。ATB的开发者在社区的讨论里提到过,他们做过实验,融合粒度不是越细越好——如果把5个以上的算子融在一起,指令缓存命中率下降,反而会变慢。所以ATB的融合规则是精心调过的,不是无脑全融。

显存带宽的影响

还有一个影响因素是显存带宽。Ascend 910的HBM带宽是1.6TB/s(理论值),看起来很大,但大模型推理的时候,不只是LayerNorm在搬数据,Attention的KV Cache、FFN的权重都在抢带宽。

ATB的LayerNorm融合省掉的HBM读写,释放出来的带宽可以给KV Cache用——这在大上下文长度(比如32K、128K)的时候效果更明显。社区里有人测过,上下文长度从512涨到8192的时候,融合带来的加速比从15%涨到了接近20%,因为这时候HBM带宽更紧张,少搬一次数据收益更大。


第四章:怎么用上ATB的LayerNorm融合

如果你现在有一个PyTorch的LLM模型,想用上ATB的融合,步骤其实不复杂。

基本流程

# step 1: 安装ATB和TorchAir
# pip install torch-air  # CANN的PyTorch适配层
# ATB一般是通过CANN统一安装的,在CANN 8.0开源版里已经包含了

import torch
import torch_npu  # 昇腾的PyTorch插件
from torch_air import TorchAir  # CANN的图优化层

# step 2: 加载你的模型(以LLaMA为例)
model = AutoModelForCausalLM.from_pretrained("your-llama-7b")

# step 3: 把模型通过TorchAir转到NPU上
# TorchAir会做子图匹配和融合,包括LayerNorm的融合
ta = TorchAir()
model_npu = ta.optimize(model)  # 这一步会自动做算子融合

# step 4: 搬到NPU上跑
model_npu = model_npu.npu()
input_ids = tokenizer("你好", return_tensors="pt")["input_ids"].npu()
output = model_npu.generate(input_ids, max_new_tokens=128)

上面的代码里,ta.optimize() 是最关键的一步。它内部会:

  1. 把PyTorch的计算图转成CANN GE图引擎的图表示
  2. 跑融合规则匹配(包括Add+LayerNorm、LayerNorm+MatMul等)
  3. 把匹配到的子图替换成融合算子
  4. 返回优化后的模型

整个过程不需要你手动修改模型代码里的LayerNorm调用。

验证融合是否生效

怎么确认LayerNorm融合真的生效了?可以用CANN的profiling工具:

# 跑推理的时候开启profiling
export ASCEND_PROFILING_MODE=1
export ASCEND_PROFILING_OPTIONS="task_trace:fp_point=Default/network-Forward/Add;bp_point=Default/network-Backward/Add"

# 跑你的推理脚本
python your_infer_script.py

# profiling结果会写到 $HOME/ascend_profiling/ 下面
# 用msprof工具查看
msprof --export=on --output=$HOME/ascend_profiling/xxx

在profiling的输出里,如果你看到 AddLayerNorm 或者 FusedAttention 这样的算子名,说明融合生效了。如果还是看到独立的 LayerNorm 算子,说明匹配没成功,可能是模型结构的问题,也可能是输入shape比较特殊(比如动态shape)导致融合规则没触发。


写在最后

LayerNorm的融合说到底不是什么黑科技,它就是软件优化里最朴素的道理:让数据少搬一次,比让计算变快一次更容易

ATB做的事情,就是把这个朴素的道理,在Transformer推理这个特定场景里做到了极致。addLayerNorm 融在一起,LayerNormMatMul 融在一起,融完之后还保证了数值正确性(ATB的CI里有一堆对齐PyTorch输出的单测),这些都是苦活累活,不是写篇论文那么简单。

仓库地址:https://atomgit.com/cann/ascend-transformer-boost

更多推荐