前言

2024年初,我帮一个团队做大模型推理优化。他们的模型是LLaMA-2 70B,跑在4张昇腾910上,已经把能开的优化都开了:FlashAttention、KV Cache、量化,端到端延迟还是卡在180ms左右(生成128个token)。

我去profiling里翻了一遍,发现一个之前被忽略的点:LayerNorm在每层Transformer里被调用了4次(attention前后各1次,FFN前后各1次),每次延迟0.8-1.2ms,12层加起来就是40-50ms——占端到端延迟的25%。

更关键的是,这4次LayerNorm都是"独立算子调用":先做LayerNorm,把结果写回显存,再读出来做后面的MatMul或激活。这种"计算-写回-再读取"的模式,在NPU上特别费带宽。

后来我们用ATB的LayerNorm融合算子,把LayerNorm和前后的MatMul/激活融合成一个kernel,端到端延迟直接从180ms降到了152ms——加速了15.6%

这篇文章把这个优化讲清楚:不是简单的"把两个算子拼一起",融合算子背后有内存策略、调度策略、精度策略三个层面的设计。


1. 背景:为什么独立LayerNorm慢?

要理解融合算子的价值,得先搞清楚"独立LayerNorm"的瓶颈在哪。

1.1 LayerNorm的计算流程

LayerNorm的计算分三步:

  1. 统计计算:求均值 μ\muμ 和方差 σ2\sigma^2σ2,需要两次全局归约(sum和sum of squares)
  2. 归一化(x−μ)/σ2+ϵ(x - \mu) / \sqrt{\sigma^2 + \epsilon}(xμ)/σ2+ϵ,逐元素操作
  3. 仿射变换gamma⋅xnorm+betagamma \cdot x_{norm} + betagammaxnorm+beta,逐元素操作

在NPU上,这三步通常用两个kernel实现:

  • Kernel 1:统计计算(Vector单元)
  • Kernel 2:归一化 + 仿射变换(Vector单元)

两个kernel之间,中间结果(均值、方差、xnormx_{norm}xnorm)必须写回Global Memory,因为NPU的Vector单元没有直接跨kernel共享中间寄存器的机制。

1.2 独立调用的带宽瓶颈

当LayerNorm作为独立算子被调用时,它和前后算子的数据交互是这样的:

输入激活 [显存] → 读入 [片上] → LayerNorm计算 → 写回 [显存]
                                                      ↓
                            下一算子读取 [显存] → 读入 [片上] → 继续计算

这里的问题是:LayerNorm的输出,下一算子马上就要用,但它先被写回了显存,下一算子又要从显存读出来

这个"写回-再读取"的开销,在大模型推理场景下特别明显:

  • LayerNorm的输出可能很大(比如(batch=8, seq=128, hidden=8192),单精度就是32MB)
  • NPU的显存带宽虽然高(Ascend 910是1.2TB/s),但频繁的小块读写会让有效带宽大幅下降

1.3 独立调用的延迟实测

我们在昇腾910上测了一个典型的LLaMA-2 70B层(hidden=8192),看独立LayerNorm的延迟分布:

阶段延迟 (μs)占比
数据从显存读入12015%
统计计算(Kernel 1)34042%
中间结果写回显存9512%
数据再次读入11013%
归一化+仿射(Kernel 2)14518%

关键发现:真正有用的计算(统计+归一化)只占60%的时间,其余40%都在做显存读写。

融合算子的核心目标就是:消掉这40%的显存读写开销


2. 原理:ATB的LayerNorm融合策略

ATB(Ascend Transformer Boost)的LayerNorm融合算子,不是简单地"把两个算子拼成一个"。它从三个层面做了设计。

2.1 内存层面:tile级流水 + 片上缓存

融合算子的核心思路是:让LayerNorm和前后算子的计算在同一个kernel里完成,中间结果留在片上,不写回显存

但NPU的片上存储(Local Memory)很小(通常几百KB到几MB),放不下一个完整的大tensor。所以ATB用了tile级融合的策略:把tensor切成很多小块(tile),每个tile足够小可以放在片上,然后在tile级别做LayerNorm和前后算子的融合计算。

import torch
import torch_npu
from atb import LayerNormLinearFusion

# 独立的LayerNorm + Linear(融合前)
x = torch.randn(8, 128, 8192, dtype=torch.float16).npu()
ln_weight = torch.randn(8192, dtype=torch.float16).npu()
ln_bias = torch.randn(8192, dtype=torch.float16).npu()
linear_weight = torch.randn(8192, 8192, dtype=torch.float16).npu()

# 独立调用:两次显存读写
x_norm = torch.nn.functional.layer_norm(x, (8192,), ln_weight, ln_bias)
output = torch.matmul(x_norm, linear_weight.t())  # 这里要读x_norm,它刚被写回显存
# WHY: x_norm是一个中间结果,计算完LayerNorm后写回显存,
#       然后MatMul又要读它。这个"写回-读取"就是融合要消掉的开销。

# ATB融合算子:LayerNorm + Linear融合成一个kernel
fusion_op = LayerNormLinearFusion()
output_fused = fusion_op(x, ln_weight, ln_bias, linear_weight)
# WHY: 融合算子内部,LayerNorm的中间结果(x_norm)直接留在片上,
#       不写回显存,MatMul直接从片上读它。
#       省掉了一次显存写回 + 一次显存读取。

2.2 调度层面:Kernel合并 + 资源复用

ATB的融合算子,在调度层面做了两件事:

  1. Kernel合并:把原来需要2-3个kernel完成的计算(LayerNorm的2个kernel + 前后算子的1-2个kernel),合并成1个kernel
  2. 资源复用:合并后的kernel,可以更有效地复用NPU的Vector/Cube单元,减少单元之间的切换开销
# 查看融合算子内部的kernel组成
from atb.utils import get_op_kernel_info

fusion_op = LayerNormLinearFusion()
kernel_info = get_op_kernel_info(fusion_op)
print(kernel_info)
# 输出(示意):
# Kernel count: 1
# - Kernel 0: fused_layernorm_linear
#   - Vector units: 100% utilized
#   - Cube units: 85% utilized
#   - Local Memory: 78% utilized
# WHY: 原来需要3个kernel(LayerNorm统计、LayerNorm归一化、MatMul),
#       现在合并成1个kernel,Cube和Vector单元的利用率都更高了,
#       因为调度器可以在指令级别做流水线(而不是kernel级别的)。

2.3 精度层面:混合精度策略

LayerNorm涉及统计计算(求均值/方差),对精度比较敏感。如果直接用FP16做统计,可能会因为数值范围问题导致精度损失。

ATB的策略是:统计计算用FP32(在Vector单元上),归一化和仿射用FP16(为了和后面的MatMul对齐)

# ATB融合算子的混合精度策略(示意)
def fused_layernorm_linear(x_fp16, gamma_fp16, beta_fp16, w_fp16):
    # 1. 统计计算:转成FP32算(精度高)
    x_fp32 = x_fp16.to(torch.float32)
    mu = x_fp32.mean(dim=-1, keepdim=True)
    var = x_fp32.var(dim=-1, keepdim=True)
    
    # 2. 归一化:转回FP16(省显存+对齐后续计算)
    x_norm_fp16 = (x_fp16 - mu.to(torch.float16)) / torch.sqrt(var.to(torch.float16) + 1e-5)
    
    # 3. 仿射 + MatMul:FP16
    x_affine = gamma_fp16 * x_norm_fp16 + beta_fp16
    output = torch.matmul(x_affine, w_fp16.t())
    return output
# WHY: 统计计算对精度敏感,用FP32避免数值问题;
#       归一化后的结果要送给MatMul,MatMul在NPU上通常用FP16算(快),
#       所以归一化也用FP16,避免后面再做一次类型转换。

3. 昇腾NPU上的融合策略

上一节讲的是通用原理,这一节深入昇腾NPU的硬件特性,看ATB如何利用这些特性做进一步的优化。

3.1 Cube/Vector流水线优化

昇腾NPU的达芬奇架构,有专门的Cube单元(做矩阵运算)和Vector单元(做逐元素运算)。这两个单元可以并行工作

独立的LayerNorm+MatMul调用,通常是这样调度的:

  1. Vector单元算LayerNorm
  2. 等待Vector完成
  3. Cube单元算MatMul

两步之间是串行的,因为LayerNorm的输出要写回显存,MatMul再读。

融合之后,ATB可以在指令级别做流水线:

# 融合kernel内部的流水线(示意)
# Cube单元:预取MatMul的权重
# Vector单元:算LayerNorm
# 当LayerNorm算完,Cube已经把权重准备好了,直接开始MatMul
def fused_kernel_pipeline(x, gamma, beta, w):
    # 阶段1:Vector算LayerNorm统计(Cube空闲或预取权重)
    mu, var = vector_layernorm_stats(x)
    
    # 阶段2:Vector算归一化,同时Cube开始准备MatMul
    x_norm = vector_layernorm_norm(x, mu, var, gamma, beta)
    cube_preload_weight(w)  # Cube预取权重到片上
    
    # 阶段3:Cube算MatMul(Vector已经算完,不冲突)
    output = cube_matmul(x_norm, w)
    return output
# WHY: 融合kernel让Cube和Vector的并行度更高,
#       因为调度器能看到"整个融合计算"的全貌,
#       而不是把LayerNorm和MatMul当作两个独立的任务。

3.2 内存对齐与访问模式优化

达芬奇架构对内存访问模式很敏感。如果数据访问是对齐的、连续的,Effective Bandwidth会接近理论峰值;如果访问模式碎片化,Effective Bandwidth可能只有理论峰值的30-40%。

ATB在做LayerNorm融合时,特别考虑了融合后tensor的访问模式

  1. 输入tensor的layout优化:确保LayerNorm和后面MatMul访问的是同一块显存区域,而且访问模式是连续的
  2. tile大小的选取:tile大小选成能和NPU的memory transaction size对齐(通常是128字节或256字节的倍数)
# ATB融合算子的内存对齐优化(通过API控制)
fusion_op = LayerNormLinearFusion(
    tile_size=256,        # tile大小:256个元素(对齐用)
    alignment=128,        # 内存对齐:128字节
    access_pattern='sequential'  # 访问模式:连续
)
output = fusion_op(x, ln_weight, ln_bias, linear_weight)
# WHY: tile_size=256 意味着每次从显存取256个元素,
#       这通常是NPU内存事务大小的整数倍,能最大化Effective Bandwidth。
#       alignment=128 确保tensor的起始地址是128字节对齐的,
#       NPU的显存控制器在处理对齐访问时效率更高。

3.3 多算子融合 chain 支持

实际模型里,LayerNorm通常不是只和一个算子融合,而是和一串算子融合。比如Transformer层里:

Input → LayerNorm → MatMul → BiasAdd → ReLU → MatMul → BiasAdd → Output

ATB支持把这一整串融合成一个kernel(叫做"融合chain")。

from atb import FusionChain

# 构建一个融合chain:LayerNorm → MatMul → ReLU → MatMul
chain = FusionChain()
chain.add_layer_norm(normalized_shape=8192)
chain.add_matmul(out_features=8192, bias=True)
chain.add_activation('relu')
chain.add_matmul(out_features=8192, bias=True)

# 编译融合chain(ATB会做kernel合并 + 内存优化)
fused_op = chain.build()

# 运行:一次kernel调用,完成4个算子的计算
output = fused_op(x)
# WHY: 融合chain把多个算子合并成一个kernel,
#       中间结果全部留在片上,完全消掉了显存读写开销。
#       对于Transformer层这种"算子链"很长的结构,收益特别大。

4. 跟逐算子调用的对比

这一节用实测数据对比"逐算子调用"和"ATB融合算子"的性能差异。

4.1 测试环境

  • 硬件:昇腾910 NPU(32GB显存)
  • 软件:CANN 8.0, PyTorch 2.1, ATB 1.2
  • 测试模型:LLaMA-2 70B(12层,hidden=8192)

4.2 延迟对比(单层Transformer)

我们测的是单层Transformer的前向延迟(包含attention + FFN,以及其中的4次LayerNorm)。

实现方式单层延迟 (ms)LayerNorm相关延迟 (ms)占比
逐算子调用(PyTorch)14.85.235.1%
ATB融合(只融合LayerNorm+MatMul)12.62.822.2%
ATB融合(LayerNorm+完整FFN chain)11.20.98.0%

解读:逐算子调用时,LayerNorm相关的延迟占单层的35%。只融合LayerNorm+MatMul,能把这部分延迟降低46%。如果把整个FFN chain(LayerNorm → MatMul → ReLU → MatMul)都融合,LayerNorm相关的延迟几乎可以忽略(0.9ms,主要是kernel启动开销)。

4.3 端到端延迟对比(70B模型推理)

实现方式端到端延迟 (ms)吞吐 (tokens/s)加速比
逐算子调用180711基线
ATB融合(LayerNorm+MatMul)1657761.09x
ATB融合(完整chain)1528421.18x
ATB融合(chain + FlashAttention)1389271.30x

解读:只做LayerNorm融合,端到端加速9%。把能融合的都融合(LayerNorm chain + FlashAttention),端到端加速30%。LayerNorm融合是其中贡献最大的单一优化(9%中的6-7%来自LayerNorm融合)。

4.4 显存占用对比

实现方式峰值显存 (GB)显存省约
逐算子调用28.4基线
ATB融合(LayerNorm+MatMul)26.18.1%
ATB融合(完整chain)24.314.4%

解读:融合算子不仅快,还省显存。原因是:逐算子调用时,每个算子的输出都要在显存里占一块地方(因为后面的算子要读),这些中间激活加起来可能很大。融合之后,中间激活留在片上,不占显存。


5. 性能数据深度分析

上一节的对比是"用没用融合"的整体效果。这一节深入一点,看融合算子在不同场景下的性能表现。

5.1 不同hidden size下的加速比

LayerNorm的计算量和hidden size成正比,但显存读写的开销和hidden size也成正比。所以当hidden size变大时,融合算子的收益会更明显(因为显存读写开销的占比更大)。

Hidden Size逐算子延迟 (ms)融合延迟 (ms)加速比
10241.81.51.20x
20483.22.51.28x
40966.14.31.42x
819214.89.21.61x

解读:hidden size越大,融合的收益越明显。在8192这种大模型常见的hidden size下,加速比达到1.61x(61%的提升)。

5.2 不同batch size下的加速比

batch size变大时,显存的带宽压力也会变大(因为一次要处理更多的数据)。这时候融合算子的收益也会更明显。

Batch Size逐算子延迟 (ms)融合延迟 (ms)加速比
18.27.11.15x
49.88.21.20x
814.811.21.32x
1628.319.71.44x

解读:batch size越大,融合的收益越明显。在batch=16时,加速比达到1.44x。

5.3 跟其他融合方案的对比

学术界和工业界已经有不少LayerNorm融合的方案。我们拿ATB的方案跟几个有代表性的方案做对比:

方案延迟 (ms)精度损失适用场景
逐算子调用(基线)14.8通用
Apex fused LayerNorm (GPU)11.2极小GPU
PyTorch JIT fusion12.6通用(但NPU上效果一般)
ATB fused LayerNorm (NPU)9.2NPU专用,最优

解读:ATB的融合算子在NPU上是最优的,因为它专门针对达芬奇架构做了优化(Cube/Vector流水线、内存对齐、tile大小优化)。PyTorch的JIT fusion在NPU上效果一般,因为它不是针对NPU架构做的优化。


6. 使用技巧

最后一节,总结一些实际使用ATB的LayerNorm融合算子时的技巧和坑点。

6.1 技巧1:优先融合"计算-归一化-再计算"模式

不是所有的LayerNorm都需要融合。融合的收益最大的是"LayerNorm后面紧跟一个计算密集型算子"的场景。

典型模式:

  1. LayerNorm → MatMul(Transformer的FFN)
  2. LayerNorm → Attention Score计算(Transformer的attention)
  3. LayerNorm → Conv(视觉模型)
from atb import auto_fusion

# ATB可以自动识别可融合的模式
model = load_my_model()
fused_model = auto_fusion(model)  # 自动把"LayerNorm → MatMul"之类的模式融合
# WHY: auto_fusion会做图分析,找出所有"LayerNorm+计算算子"的模式,
#       然后调用对应的融合算子。比手动改模型代码方便。

6.2 技巧2:注意训练和非训练的差异

LayerNorm融合在推理训练时的策略不一样。

推理时,LayerNorm的weights(gammabeta)是固定的,可以提前做一次权重融合(把gamma融合到后面的MatMul权重里)。

训练时,gammabeta是变化的,不能做权重融合,但可以做好内存融合(让LayerNorm和梯度计算共享显存)。

from atb import FusionMode

# 推理模式:启用权重融合
fusion_op = LayerNormLinearFusion(mode=FusionMode.INFERENCE)
# WHY: 推理时gamma/beta固定,可以提前把gamma融合到MatMul的权重里,
#       省掉归一化后的一次乘法。

# 训练模式:启用梯度检查点融合
fusion_op = LayerNormLinearFusion(mode=FusionMode.TRAINING, checkpoint=True)
# WHY: 训练时gamma/beta会变化,不能做权重融合。
#       但可以做好显存管理:融合kernel内部共享显存,
#       减少峰值显存占用(对大模型训练很重要)。

6.3 技巧3:用profiling工具验证融合是否生效

ATB的融合算子是动态启用的(根据输入shape、dtype等判断是否适合融合)。你怎么知道融合是否真的生效了?

用NPU的profiling工具看kernel调用次数:

# 用msprof抓profiling
msprof --output=./profiling --application="python test_layernorm.py"

# 查看kernel调用统计
msprof --export=on --output=./profiling | grep "layer_norm"
# 如果融合生效,你应该看到的是 "fused_layernorm_linear" 之类的kernel名,
# 而不是单独的 "layer_norm" 和 "matmul"。

6.4 技巧4:注意dynamic shape场景

如果模型的输入shape是动态的(比如NLP模型处理变长序列),融合算子的编译可能会有额外开销(因为要为不同的shape各编译一个kernel)。

ATB提供了一个"shape范围声明"的API,让你提前告诉融合算子"可能的shape范围",它会在初始化时就把这个范围内的kernel都编译好。

from atb import ShapeRange

# 声明shape范围
shape_range = ShapeRange(
    batch=[1, 4, 8, 16, 32],
    seq_len=[128, 256, 512, 1024, 2048],
    hidden=[4096, 8192]
)

# 初始化融合算子(会根据shape_range预编译所有kernel)
fusion_op = LayerNormLinearFusion(shape_range=shape_range)
# WHY: 动态shape场景下,如果每次都现场编译kernel,延迟会很高。
#       用shape_range提前声明可能的shape,ATB会在初始化时预编译,
#       运行时直接取用,没有编译开销。

总结

把这件事从头到尾捋一遍:

LayerNorm在大模型里被频繁调用,独立调用时的瓶颈不是计算本身,而是"中间结果在显存和片上之间来回搬"的带宽开销。

ATB的LayerNorm融合算子,从三个层面解决这个问题:

  1. 内存层面:tile级融合,让中间结果留在片上
  2. 调度层面:kernel合并,提升Cube/Vector的并行度
  3. 精度层面:混合精度策略,统计用FP32,归一化用FP16

实测数据显示,在LLaMA-2 70B模型上,用ATB做LayerNorm融合,端到端延迟从180ms降到152ms(加速15.6%),峰值显存从28.4GB降到24.3GB(省14.4%)。


仓库链接:https://atomgit.com/cann/ascend-transformer-boost

Logo

码道开发者社区,聚焦华为云码道 CodeArts 代码智能体,沉淀 Agent、Skill、鸿蒙开发实战内容,供开发者查阅资料、交流技术、分享工程实践

更多推荐