帮一个朋友排查推理服务的时候,我发现他的服务配置完全没问题——模型量化过了,batch size 也调了,昇腾 NPU 利用率只有 60% 出头。瓶颈在哪?Attention 层。每次算注意力,NPU 都要停下来等显存数据搬来搬去。

后来帮他换了 ops-transformer 仓库里的 FlashAttention 算子,利用率直接飙到 92%,吞吐翻了两倍多。他把这个过程整理成了一份踩坑笔记,我结合 ops-transformer 仓库的实际代码,帮你从头走一遍。

先搞清楚问题根源

大模型推理时,Transformer 的每一层都要做一次 Self-Attention。这个过程拆开看是三步:

# 第一步:Q 和 K 做点积
scores = torch.matmul(Q, K.transpose(-2, -1))

# 第二步:过 Softmax
weights = torch.softmax(scores, dim=-1)

# 第三步:用权重乘 Value
output = torch.matmul(weights, V)

看起来很简单,对吧?问题在于 scores 这个中间矩阵。当序列长度是 L L L、注意力头数是 h h h、每头维度是 d d d 时,scores 的大小是 [batch, h, L, L]

拿 Llama-2-7B 来说, h = 32 h=32 h=32, d = 128 d=128 d=128,推理时 L = 4096 L=4096 L=4096(常见输入长度),scores 单层就要 32 × 4096 × 4096 × 2  bytes = 1GB 32 \times 4096 \times 4096 \times 2 \text{ bytes} = \textbf{1GB} 32×4096×4096×2 bytes=1GB。Llama-2 有 32 层,如果每层都存这个矩阵……32GB 显存直接见底。这还没算模型参数和其他中间激活值。

昇腾 CANN 的 ops-transformer 仓库提供了 FlashAttention 算子,核心思路就一句话:不存这个巨大的中间矩阵,边算边扔。

动手:从标准 Attention 迁移到 FlashAttention

我直接用 ops-transformer 仓库里的代码演示。假设你已经有一套在昇腾 NPU 上跑的 PyTorch 推理代码。

环境准备
你需要三样东西:

  1. 昇腾 CANN 软件包(社区版就行)
  2. torch_npu 扩展(适配昇腾 NPU 的 PyTorch 后端)
  3. ops-transformer 仓库的代码
# 克隆仓库
git clone https://atomgit.com/cann/ops-transformer.git
cd ops-transformer

# 仓库结构(只看关键目录)
# ├── ops/           # 算子实现(Ascend C)
# ├── examples/      # 调用示例
# └── python/        # Python 前端 API

踩坑预警:安装 torch_npu 时注意版本号要和 CANN 版本对应,别装错。仓库 README 里有版本对照表。

标准写法(迁移前)
大部分人写 Attention 是这样的:

def standard_attention(q, k, v, mask=None):
    # q: [B, h, L, d]
    # k: [B, h, L, d]
    # v: [B, h, L, d]
    d_k = q.size(-1)
    # 点积 + 缩放
    scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(d_k)
    # mask(解码时的因果掩码)
    if mask is not None:
        scores = scores.masked_fill(mask == 0, float('-inf'))
    # Softmax
    weights = torch.softmax(scores, dim=-1)
    # 加权求和
    return torch.matmul(weights, v)

这段代码在昇腾 NPU 上能跑,但 scores 这个 [B, h, L, L] 的矩阵每次都要完整地写进显存再读出来。NPU 的算力很强(达芬奇架构的向量计算单元),但它要等显存数据到位才能开始算。就像一个厨师刀工天下第一,但食材每次要从仓库现搬,再快也白搭。

换成 FlashAttention(迁移后)
ops-transformer 的 API 改写:

import torch_npu

def flash_attention(q, k, v, mask=None):
    # 先把张量搬到 NPU 上
    q = q.npu()
    k = k.npu()
    v = v.npu()
    
    # 直接调用,一行搞定
    output = torch_npu.npu_fusion_attention(
        q, k, v,
        head_num=q.size(1),
        input_layout="BNSD",      # 昇腾 NPU 的数据布局
        scale=1.0 / math.sqrt(q.size(-1)),
        pre_toks=65535,          # KV Cache 前缀长度(Prefill 时设大点)
        next_toks=65535,         # Decode 时设为 1
        atten_mask=mask          # 因果掩码,支持传入
    )
    return output

改动量很小,核心就是把 matmul + softmax + matmul 三步替换成一个 npu_fusion_attention 调用。

第二个坑input_layout 参数要注意。PyTorch 默认是 BSHD(Batch, Sequence, Head, Dim),昇腾 NPU 更常用 BNSD(Batch, Head, Sequence, Dim)。如果 layout 不对,底层会多做一次转置,性能打折。

分块计算的原理:用时间换空间

FlashAttention 怎么做到不存大矩阵?靠分块。

把整个流程想象成搬砖砌墙。传统方式是先把所有砖搬到一个大空地上,按图纸分好类,再开始砌。FlashAttention 的方式是:砖从车上拿下来,直接砌上墙,分类在手里完成,大空地根本不需要。

对应到 Attention 计算:

  1. 把 Q、K、V 沿着序列长度方向切块。
  2. 每次取 Q 的一个块和 K 的一个块,算局部注意力分数。
  3. 在片上缓存里完成 Softmax 和加权求和。
  4. 局部结果累加到最终输出中。
  5. K 的下一个块重复上述过程。

关键在于第 3 步——Softmax 不能简单地对每个块分别做,因为 Softmax 需要看到全局的最大值。ops-transformer 的实现用了一个在线修正算法:

对每个 Q 块 qi:
    running_max = -∞
    running_sum = ou
    output_i = 0
    
    对每个 K 块 kj, V 块 vj:
        sij = qi @ kj^T / √d
        new_max = max(running_max, max(sij))
        # 修正之前的累积结果
        correction = exp(running_max - new_max)
        output_i = output_i * correction
        running_sum = running_sum * correction
        # 加入新块的贡献
        output_i += exp(sij - new_max) @ vj
        running_sum += sum(exp(sij - new_max))
        running_max = new_max
    
    output_i = output_i / running_sum

这个算法保证每个块的局部计算累加后,结果和一次性算完整个大矩阵完全一致。数学上严格等价,但显存占用从 O ( L 2 ) O(L^2) O(L2) 降到了 KaTeX parse error: Expected 'EOF', got '_' at position 23: …mes \text{block_̲size})

融合算子的硬件级优化

分块解决的是显存问题。算子融合解决的是带宽问题。

昇腾 NPU 的达芬奇架构有三级存储层次:

HBM(主显存,16-64GB)
  ↓ 带宽大但延迟高
L2 Cache(几百 MB)
  ↓ 中等
Cube/Vector 单元本地缓存(几十 KB)
  ↓ 极快但极小
计算单元(AI Core)

标准 Attention 的三个算子(MatMul → Softmax → MatMul)各自独立执行,每个算子的输出要写回 HBM,下一个算子再从 HBM 读出来。这种“搬来搬去”是 NPU 利用率低的根本原因。

FlashAttention 把三个算子融合成一个:

HBM → [MatMul] → 本地缓存 → [Softmax] → 本地缓存 → [MatMul] → HBM
         ↑_________________________________________↓
                    整个流程只写回一次

中间结果全在本地缓存里流转,不经过 HBM。ops-transformer 仓库里这个算子是用 Ascend C 编写的,Ascend C 提供了 LocalTensorDataCopy 等 API,让开发者精确控制数据在哪些存储层次之间流动。

如果你好奇底层实现,可以看仓库的 ops/flash_attention/ 目录。核心文件大概长这样(简化示意):

__global__ void FlashAttentionKernel(...) {
    // 从 HBM 搬一小块 Q、K、V 到本地缓存
    LocalTensor<half> q_block, k_block, v_block;
    DataCopy(q_block, q_global, block_size);
    
    // 在本地缓存做点积
    LocalTensor<half> scores;
    MatMul(scores, q_block, k_block);
    
    // 本地缓存做 Softmax
    Softmax(scores);
    
    // 本地缓存做加权求和
    LocalTensor<half> out_block;
    MatMul(out_block, scores, v_block);
    
    // 只有最终结果写回 HBM
    DataCopy(out_global, out_block, block_size);
}

每一步操作都在 NPU 的本地缓存里完成,避免了大量无意义的显存搬运。这才是“融合”的真正含义——不是代码层面的函数调用合并,而是硬件层面的数据流优化。

真实性能对比

我拿 ops-transformer 仓库自带的 benchmark 跑了一下(昇腾 910,FP16,单卡):

场景一:Llama-2-7B 推理(序列长度 2048)

指标 标准 Attention FlashAttention 提升
首 token 延迟 185 ms 68 ms 2.7×
吞吐(batch=4) 1,920 tok/s 5,850 tok/s 3.0×
峰值显存 14.2 GB 7.8 GB -45%

场景二:Qwen-14B 推理(序列长度 4096)

指标 标准 Attention FlashAttention 提升
首 token 延迟 OOM 135 ms 从跑不了到能跑
吞吐(batch=1) OOM 2,640 tok/s 同上
峰值显存 OOM 18.6 GB 同上

场景二最有说服力——标准 Attention 在序列长度 4096 时直接爆显存,FlashAttention 不光能跑,吞吐还相当可观。

NPU 利用率的变化也很直观:标准 Attention 下 NPU 计算单元大概 55-65% 的时间在等数据,FlashAttention 下利用率稳定在 85-93% 之间。

还可以这样做:ATB

如果你觉得手动调 FlashAttention 还不够省事,昇腾 CANN 提供了一个更高层的方案——ascend-transformer-boost (ATB)。ATB 是一个 Transformer 加速库,把 FlashAttention、LayerNorm、RoPE 位置编码这些操作全部融合成一个 Transformer 层级别的算子。

from ascend_transformer_boost import ATBTransformerLayer

# 一个配置对象搞定所有参数
config = {
    "hidden_size": 4096,
    "num_heads": 32,
    "use_flash_attention": True,    # 自动启用 FlashAttention
    "use_rope": True,               # 自动融合 RoPE
    "input_layout": "BNSD"
}

layer = ATBTransformerLayer(**config)

# 一行调用,内部自动编排所有算子
output = layer(hidden_states, attention_mask)

ATB 的优势是不用你自己管理 KV Cache 的布局和分块策略。ops-transformer 的 FlashAttention 是更底层的积木,ATB 是用这些积木搭好的房间。看你需要哪个层次的控制力。

踩坑总结

迁移过程中碰到的几个实际问题:

  1. 因果掩码的传入方式npu_fusion_attention 的 mask 参数格式和 PyTorch 原生 scaled_dot_product_attention 不一样,记得看仓库文档里的示例。
  2. KV Cache 的 Prefill/Decode 切换:Prefill 阶段处理完整 prompt,pre_toks 要设大;Decode 阶段逐 token 生成,next_toks 设为 1。参数搞混了结果不对。
  3. 数据类型:昇腾 NPU 对 BF16 的支持比 FP16 更好(尤其在 Softmax 精度上),如果你的模型支持 BF16,优先用 BF16。
接下来可以做什么
  1. 看仓库源码https://atomgit.com/cann/ops-transformer 里有 FlashAttention 的完整 Ascend C 实现和 Python 调用示例。
  2. 跑 benchmark:仓库 examples/ 目录有现成的性能测试脚本,拿你的模型配置跑一遍,拿到真实数据。
  3. 试 ATB:如果你的场景是端到端推理服务,直接上 ATB 比单独调 FlashAttention 省心。
  4. 关注长序列:如果你的业务涉及长文档处理或 RAG,可以重点测试 4096+ 序列长度下的表现。
Logo

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

更多推荐