LayerNorm也能融合?ATB让你的大模型推理再快15%
前言
去年做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上的执行方式是这样的:
- Vector单元一边做加法,一边累积均值和方差需要的统计值(部分求和、部分平方和)
- 加法做完的同时,均值和方差也算好了
- 直接用算好的均值方差做归一化,不需要再把数据读一遍
这样做的好处是:原来需要两次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() 是最关键的一步。它内部会:
- 把PyTorch的计算图转成CANN GE图引擎的图表示
- 跑融合规则匹配(包括Add+LayerNorm、LayerNorm+MatMul等)
- 把匹配到的子图替换成融合算子
- 返回优化后的模型
整个过程不需要你手动修改模型代码里的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推理这个特定场景里做到了极致。add 和 LayerNorm 融在一起,LayerNorm 和 MatMul 融在一起,融完之后还保证了数值正确性(ATB的CI里有一堆对齐PyTorch输出的单测),这些都是苦活累活,不是写篇论文那么简单。
仓库地址:https://atomgit.com/cann/ascend-transformer-boost
更多推荐

所有评论(0)