刚接触大模型那会儿,我被 Attention 机制的显存占用砸懵了。一个 70 亿参数的模型,推理时光 Attention 就要吃掉几十 GB 显存,序列长度一上来,直接爆卡。

后来才发现,问题不在于模型本身,而在于 Attention 的计算方式。

昇腾 CANN 的 ops-transformer 仓库里有个 FlashAttention 算子,专门解决这个痛点。它在昇腾 NPU 上通过算子融合和显存优化,把大模型推理速度提了 3 倍左右,显存占用还能砍掉一半。

今天就拆解一下这个算子到底做了什么。

1. 先搞清楚 Attention 的显存坑在哪

标准的 Attention 计算大家都熟悉:Query、Key、Value 三个矩阵做点积,过 Softmax,再跟 Value 相乘。

问题出在中间那个 Softmax——它会生成一个巨大的注意力矩阵,维度是 [序列长度 × 序列长度]

  • 举个具体的例子:假设序列长度是 4096(在长文本场景很常见),那注意力矩阵就是 4096 × 4096 = 16 M 4096 \times 4096 = 16M 4096×4096=16M 个元素。
  • 如果用 FP16 存储,一个元素 2 字节,单层的注意力矩阵就要 32MB 显存。
  • 多层 Transformer 叠起来,显存直接炸掉。

更要命的是,这个注意力矩阵只是中间结果,算完就扔,真正需要的只是最终输出。但你必须在显存里把它存下来,因为反向传播时要算梯度。

这就是所谓的 “显存墙”——计算量没多少,但显存先扛不住了。

FlashAttention 的核心思路:不存这个大矩阵,算完直接扔。

2. 分块计算:把大矩阵切成小块

FlashAttention 的第一个优化是 分块计算(Tiling)。昇腾 NPU 上有高速缓存,但容量有限。FlashAttention 把 Query、Key、Value 都切成小块,每次只让一小块数据进缓存算完,再换下一块。

  • 打个比方:你要把 1000 页的文档录入电脑,一次全摊开肯定摆不下。不如每次拿 10 页,录完换下 10 页。虽然多跑几趟,但不用找个篮球场那么大的桌子。

具体到 FlashAttention 的实现:

  1. 把 Query 分成多个小块,每块大小 [block_size, head_dim]
  2. Key 和 Value 也分块,大小 [block_size, head_dim]
  3. 每次拿一个 Query 块,跟所有 Key 块算注意力分数。
  4. 算完 Softmax 后,马上跟对应的 Value 块乘起来,得到这部分输出。
  5. 把结果累加到最终输出里,中间的注意力分数矩阵就可以扔了。

这个过程中,注意力矩阵只在缓存里存在一小会儿,算完立刻释放,根本不用进主显存。

3. 算子融合:减少数据搬运

FlashAttention 的第二个优化是 算子融合。昇腾 NPU 的达芬奇架构支持灵活的向量计算,FlashAttention 把 MatMulSoftmax、再 MatMul 这三个操作融合成一个算子,中间结果不用写回显存,直接在片上缓存流转。

  • 传统方式是这样的

    显存 → MatMul → 显存 → Softmax → 显存 → MatMul → 显存
    

    三次读写显存。

  • FlashAttention 是这样

    显存 → [MatMul + Softmax + MatMul] → 显存
    

    只读写一次显存。

昇腾 NPU 的内存带宽本来就不算瓶颈,但能省的地方当然要省。融合算子把中间结果的显存带宽全砍掉了,推理速度自然上去。

ops-transformer 仓库里,FlashAttention 算子是用 Ascend C 语言编写的。Ascend C 允许开发者直接控制 NPU 的计算单元和缓存,实现这种细粒度的融合优化。

4. Softmax 的数值稳定性

还有一个技术细节:分块 Softmax 的数值稳定性。

Softmax 计算时要做指数运算,数值太大容易溢出。标准做法是先减去最大值,再做 exp。但在分块计算时,每个块的最大值不一样,直接算会出错。

FlashAttention 用了一个叫 “在线 Softmax” 的技巧:

  1. 每算完一个块,记录当前的最大值和指数和。
  2. 下一个块算完后,用新的最大值修正之前的计算结果。
  3. 逐步累积,最终得到正确的 Softmax 输出。

这个修正过程需要额外的减法和乘法,但相比节省的显存带宽,这点计算开销可以忽略不计。

5. 实测性能:昇腾NPU 上的表现

根据 ops-transformer 社区仓库的最新基准测试数据(基于 2026 年 5 月的 CANN 8.0 版本),在昇腾 910 上跑 Llama2-70B 模型的实测数据如下:

配置 序列长度 吞吐 (tokens/s) 显存占用 (GB)
标准 Attention 2048 1,850 12.4
FlashAttention 2048 5,420 6.2
标准 Attention 4096 OOM -
FlashAttention 4096 3,180 11.8

可以看到:

  1. 序列长度 2048 时:吞吐提升 2.9 倍,显存减半。
  2. 序列长度 4096 时:标准 Attention 直接爆显存,FlashAttention 还能跑。

这个提升主要来自两方面:

  • 显存带宽减少:融合算子把中间结果的读写砍掉了。
  • 缓存命中率提高:分块计算让数据尽可能留在片上缓存。
6. 在 ops-transformer 里怎么用

ops-transformer 仓库里的 FlashAttention 算子已经封装好了,直接调用就行。

  • 如果你用 PyTorch 框架,可以通过 torch_npu 扩展来用:

    import torch_npu
    
    # 假设你有 Q, K, V 三个张量
    # flash_attention 会自动做分块和融合
    output = torch_npu.npu_fusion_attention(
        query,  # [batch, seq_len, num_heads, head_dim]
        key,
        value,
        head_num=num_heads,
        input_layout="BNSD",  # 昇腾NPU 常用布局
        scale=1.0 / math.sqrt(head_dim)
    )
    

    底层实现会自动选择最优的分块大小,根据序列长度和头数调整。你不用关心具体细节,只要确保输入张量已经在昇腾NPU 上(.npu())。

  • 通过 ATB 加速库调用
    昇腾 CANN 还提供了一个更上层的 API,通过 ascend-transformer-boost(ATB)加速库调用。ATB 把 FlashAttention 封装成了 Transformer 层级别的组件,连 Position Embedding 和 LayerNorm 都一起优化了:

    from ascend_transformer_boost import ATBTransformerLayer
    
    # 一行代码创建优化后的 Transformer 层
    layer = ATBTransformerLayer(hidden_size, num_heads, ...)
    
    # 内部已经用了 FlashAttention + 融合 LayerNorm
    output = layer(hidden_states)
    

    ATB 的优势是不用自己拼算子,但如果你只需要单独的 FlashAttention,ops-transformer 的 API 更轻量。

7. 跟其他方案有啥区别

技术上,昇腾的 FlashAttention 和 NVIDIA 的核心思路一样:分块计算 + 算子融合。区别在于底层实现:

  1. 硬件架构不同:NVIDIA GPU 用 CUDA,昇腾 NPU 用 Ascend C。昇腾的达芬奇架构有专门的向量计算单元,Ascend C 可以直接编程控制这些单元,实现细粒度的算子融合。
  2. 分块策略不同:昇腾 NPU 的缓存大小和 GPU 不一样,FlashAttention 的分块参数要针对昇腾架构调优。ops-transformer 仓库里的实现已经针对 Ascend 910/950 做了适配。
  3. 框架集成不同:NVIDIA 的 FlashAttention 主要通过 flash-attn Python 包调用,昇腾的是通过 torch_npu 或 ATB 调用。如果你有现成的 PyTorch 代码,迁移成本就是改几个 import。

性能上,两边都能把大模型推理速度提 2-3 倍。昇腾 NPU 在长序列场景(4096+)的优势更明显,因为显存带宽优化做得更激进。

8. 实际应用场景

FlashAttention 在这些场景最有价值:

  1. 长文本推理:RAG(检索增强生成)场景下,上下文动辄几万 token。没有 FlashAttention,显存直接炸。有了之后,序列长度 8192 也能稳稳跑。
  2. 大批量推理:并发请求多的时候,batch size 上去,Attention 的显存占用跟着涨。FlashAttention 让你在同样的显存下,能跑更大的 batch。
  3. 推理延迟优化:第一 token 延迟(TTFT)是用户体验的关键指标。FlashAttention 减少了显存访问次数,首 token 出来的更快。
    ops-transformer 仓库的完整代码在这里:
    https://atomgit.com/cann/ops-transformer

有 FlashAttention 的详细文档和性能测试脚本,直接 clone 下来就能跑

Logo

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

更多推荐