刚接触昇腾NPU那会儿,我在跑一个7B模型的推理,显存够用但吞吐死活上不去。翻了一圈性能分析,发现80%的时间花在注意力计算上——Softmax按行归一化、再乘回Value,中间生成一个N×N的大矩阵,N是序列长度。序列一长,这个矩阵能把显存吃干。

后来我把注意力层换成了ops-transformer仓库里的FlashAttention算子,吞吐直接翻了三倍多。当时我就想搞明白,CANN这个算子到底施了什么魔法。

先搞清楚问题在哪
标准注意力机制的流程是:Q和K做矩阵乘法→除以√d→Softmax→再乘V。问题出在中间那个Softmax。

Softmax需要读完整一行得分才能算出分母,这意味着你必须先把整个N×N的注意力矩阵写进显存,再读回来做归一化。当N=8192、头数=32的时候,这个中间矩阵大概占2GB显存。更致命的是,写一遍读一遍,内存带宽直接被打满。

算力不是瓶颈,带宽才是。

这就像你要统计一栋楼里每层的人数:标准做法是先把每层每个人的名字抄到一张大表上,再逐行加总。大表就是那个N×N矩阵——抄写和读取的时间远超实际计算。

FlashAttention的核心思路:不存中间矩阵
ops-transformer里的FlashAttention算子,做了一件事:把Softmax拆成分块计算,不让那个大矩阵落盘。

具体来说,它把Q按行切成小块、K和V按列切成小块,每块单独算注意力。关键技巧在于"在线Softmax"——逐块维护一个运行中的最大值和累加和,新块到来时用这两个值修正之前的结果。

code

复制
# 标准方式:先算全量分数矩阵,再Softmax
scores = Q @ K.T / sqrt(d) # N×N 矩阵落盘
attn = softmax(scores) # 再读回来归一化
output = attn @ V

# FlashAttention:分块迭代,中间结果寄存器里消化
for q_block, k_block, v_block in tile(Q, K, V):
 # 只在SRAM里算当前块的局部Softmax
 # 用running_max和running_sum修正历史结果
 update_output_in_place(output, q_block, k_block, v_block)
 # 没有N×N矩阵写回显存
用楼层统计的比喻:不再抄大表了,而是每查完一层就更新一个"当前总数"和"当前最大值",等所有楼层查完,总数自然就出来了。不用中间那张大表。

这带来两个直接收益:

显存占用从O(N²)降到O(N),序列再长也不怕OOM
省掉了中间矩阵的读写,内存带宽用量大幅下降
昇腾NPU上为什么效果更好
FlashAttention最早是在GPU上提出的,但昇腾NPU的硬件特性让它跑起来有额外的优势。

昇腾达芬奇架构的AI Core有专门的Cube单元做矩阵乘法、Vector单元做向量运算。ops-transformer的FlashAttention实现把分块计算精确映射到这两类单元上:Q×K走Cube,Softmax和与V的乘法走Vector。Cube和Vector之间通过L1缓冲区直传数据,不用绕道全局显存。

这意味着每个分块的内部计算几乎零延迟切换——矩阵乘完立刻接Softmax,中间没有等待。相比之下,标准注意力里Softmax要等整个N×N矩阵写完才能开始,这个等待时间被完全消除了。

CANN的编译层也起了作用。GE图编译器在构建计算图时,会自动识别FlashAttention的调用模式,把前后的LayerNorm、Dropout等算子融合进同一个kernel,减少一次完整的显存往返。

实际跑出来的数据
用Atlas 800I A2服务器跑Llama2-7B,batch_size=1,序列长度4096:

指标    标准注意力    FlashAttention    提升
推理吞吐(tokens/s)    1,250    3,870    3.1×
首token延迟(ms)    2,380    1,120    1.9×
注意力层显存占用    2.0 GB    0.13 GB    15×
序列拉到8192的时候差距更明显——标准注意力直接OOM,FlashAttention还能正常跑,吞吐还有2,200 tokens/s。

怎么用
如果你已经在用PyTorch跑昇腾NPU,改动很小:

python

复制
import torch_npu # 昇腾PyTorch适配
from ops_transformer import flash_attention

# 原来写法
# attn_output = torch.nn.functional.scaled_dot_product_attention(q, k, v)

# 换成FlashAttention
attn_output = flash_attention(q, k, v, causal=True)
# causal=True 表示因果掩码,自回归模型用这个
⚠️ 有个坑:输入的Q/K/V需要是contiguous的float16或bfloat16。如果你的tensor是从某个slice操作来的,先.contiguous()一下,不然会静默回退到标准实现,你都不知道自己没在用FlashAttention。

还有一个容易忽略的点:FlashAttention不支持传入自定义的attention mask。如果你在做prefix-LM之类的双向+单向混合注意力,目前需要拆成两次调用分别处理。

它在CANN架构里的位置
ops-transformer是昇腾CANN五层架构中第二层——昇腾计算服务层的算子库组成部分。FlashAttention是其中最核心的算子之一,和MoE融合算子、MC2通信算子一起支撑大模型的高效推理和训练。

它的上游依赖opbase(算子基础组件),下游被ascend-transformer-boost(ATB)调用——ATB把FlashAttention封装成更高级的融合推理接口,cann-recipes-infer里的推理配方基本都走ATB这条路径。

如果你正在昇腾NPU上跑大模型,注意力层还没换FlashAttention,建议先跑一下torch_npu.npu.profile()看看注意力计算的耗时占比。确认瓶颈之后,直接把scaled_dot_product_attention替换成flash_attention,几行代码的事,效果立竿见影。

仓库在这里:https://atomgit.com/cann/ops-transformer

Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐