前言

上个月有个群友在 CANN 社区问:“ops-transformer 到底是个算子库还是一个框架?我看它里面有 FlashAttention,也有 MoE,还有一大堆我看不懂的调度逻辑——它到底是一层还是多层?”

这个问题问得好。很多人 clone 了仓库,对着 src/ 目录看了半天,还是搞不清这个库在 CANN 五层架构里到底处于什么位置、FlashAttention 又是怎么从一堆 Ascend C 代码变成你在 Python 里调用的一行 API 的。

所以这篇文章不写教程、不写实战,就拆一件事:ops-transformer 的架构是怎么设计的,FlashAttention 在里头处于什么位置,数据从进来到出去经过了哪些层。


一、ops-transformer 不是"一层",是一组"能力"

先纠偏一个认知:ops-transformer 不是一个层级,是一个算子集合的仓库。

CANN 五层架构是这么分的:

第1层:AscendCL(编程接口层)
第2层:AOL算子库(ops-transformer 属于这一层)
第3层:Graph Compiler / ATC(编译层)
第4层:Runtime / Graph Executor(执行层)
第5层:驱动 / 基础层

ops-transformer 的位置在第 2 层,但它的内部架构并不是扁平的"一堆算子放在一起"——它有自己的三层子架构

ops-transformer 内部架构
├── 接口层(Interface Layer)
│ └── 对接 AscendCL / ATB 的调用接口
├── 算子实现层(Kernel Layer)
│ └── FlashAttention / MoE / MC2 等算子的 Ascend C 实现
└── 调度优化层(Schedule Layer)
 └── Tiling 策略、融合决策、多核并行调度

认知纠偏: 你 import 的 atb.flash_attention 不是直接调用 flash_attention_kernel.cpp——中间经过了接口层做参数校验和格式转换,再经过调度优化层决定 tiling 参数和多核分配方式,最后才进到算子实现层真正算。

这三层是串联的,不是并联。数据必须逐层往下走。


二、FlashAttention 在 ops-transformer 里的完整路径

把 FlashAttention 的调用链路拆开,从 Python API 到 NPU 真正执行,一共经过 5 个阶段

阶段 1:Python 接口层(atb / pytorch 适配)

from atb import flash_attention
output = flash_attention(query, key, value, config)

这层做的事情很简单:把 Python 对象转成 C 接口能认的结构体

// interface/flash_attention_interface.cpp
atbStatus_t FlashAttentionInference(...)
{
 // 第一步:参数校验
 if (query.dim() != 4) return ATB_ERROR;
 // 第二步:转成内部 Tensor 结构体
 atbTensor_t q_tensor = ConvertToATBTensor(query);
 // 第三步:下发到调度层
 return ScheduleFlashAttention(q_tensor, ...);
}

这里容易出的 bug: query/key/value 的 head_dim 不一致时,这层会报一个很含糊的错误 "ATB_ERROR_INVALID_PARAM"——不会告诉你哪个维度不对,得自己打日志查。

阶段 2:调度优化层(Tiling + 多核分配)

这是 ops-transformer 里最复杂的一层,也是 FlashAttention 性能的关键。

核心决策有两个:

① Tiling 参数怎么定?

输入:batch=1, heads=12, seq_len=2048, head_dim=64
Ub 容量:64KB / Core

决策逻辑:
 - head_dim=64 → 一个 token 的 Q 向量占 64×2byte(FP16) = 128B
 - Ub 能放多少个 token 的 QK^T 中间结果?
 - 算出 tile_size=128 → 刚好能塞进 Ub,不溢出到 HBM

② 多核怎么分配?

昇腾 910 有 32 个 AI Core。FlashAttention 的并行策略是:batch × heads 维度并行,seq_len 维度 tiling

// schedule/flash_attention_schedule.cpp
void ScheduleFlashAttention(...)
{
 int cores_needed = batch * heads; // 12 个 head → 至少占 12 个 Core
 int tiles_per_head = seq_len / TILE_SIZE; // 2048/128 = 16 个 tile
 // 把 16 个 tile 分配到剩余 Core 上并行算
}

阶段 3:算子实现层(Ascend C Kernel)

到这层就是纯计算逻辑了,用 Ascend C 写,跑在 NPU 的 AI Core 上。

// kernel/flash_attention_kernel.cpp
__global__ void FlashAttentionKernel(...)
{
 __shared__ float attn_scores[TILE_SIZE][TILE_SIZE]; // 放 Ub,不写回 HBM

 // ① QK^T(矩阵乘)
 MatMul(Q_tile, K_tile, attn_scores);

 // ② 在线 Softmax(一遍扫描,不回 HBM)
 OnlineSoftmax(attn_scores, ...);

 // ③ Attn × V(矩阵乘,结果还在 Ub)
 MatMul(attn_scores, V_tile, output_tile);

 // ④ 只有这步才写回 HBM
 WriteOutputToHBM(output_tile, ...);
}

加粗金句:FlashAttention 快的真正原因不是"算得更快",是"中间结果不出来,最后一步才回 HBM"。

阶段 4:融合决策(可选,由 ATB 触发)

如果你用的是 ATB 的融合接口(atb.flash_attention_fused),在阶段 2 和阶段 3 之间还会插入一个融合决策步骤

标准流程:QK^T → Softmax → AV(3 次 kernel launch)
融合流程:QK^T + Softmax + AV 融合成 1 个 kernel(1 次 kernel launch)

融合的决策在编译期就定好了(ATC 编译器做图融合),运行时直接调融合后的 kernel,少两次 launch 开销。

阶段 5:执行层(Runtime 下发到 NPU)

融合 kernel 编译好后,通过 Runtime 下发到 NPU 执行:

// runtime 下发(这层不用你管,Runtime 自动做)
runtime.EnqueueKernel(fused_kernel, ...);

三、为什么这套架构能打?

拆完架构,回过头来看为什么 ops-transformer 的 FlashAttention 在昇腾 NPU 上能比标准实现快 3 倍

答案在架构的三个设计决策里:

决策 1:Tiling 参数硬编码为 128×128(FP16)

不是动态算的,是按 Ub 容量手动调优后写死在代码里的。动态算 tiling 多一次 CPU 侧计算,静态写死零开销。代价是换其他精度(比如 FP32)要手动改参数——但昇腾上大模型推理 99% 场景用 FP16,写死没问题。

决策 2:在线 Softmax 放在算子实现层,不暴露给上层

标准 Softmax 要两遍扫描(找 max → 算 exp → 归一化),在线 Softmax 一遍搞定。但这个算法依赖 Ub 容量足够放下整个 tile 的中间结果——如果 Tiling 参数算错了(tile 太大,Ub 放不下),在线 Softmax 会直接算错,不报异常。

这就是很多人"跑通了但结果不对"的根源:TILE_SIZE 跟 Ub 容量不匹配,在线 Softmax 数值溢出。

决策 3:融合决策在编译期做,不在运行期

ATC 编译器在编译图的时候就把"能融合的算子"标记出来,运行时直接调融合 kernel。如果融合决策放在运行期做(比如每轮 inference 之前判断一次"这次能不能融合"),每次多 5-10ms 的判断开销,长序列推理拖不起。


四、跟 ascend-transformer-boost(ATB)的架构关系

最后把 ops-transformer 和 ATB 的关系在架构层面说清楚:

你的 Python 代码
 ↓
ATB(加速库层,第 2 层)
 ├── 接口封装:flash_attention()
 ├── 融合决策:编译期决定要不要融合
 └── 下发到 ↓
ops-transformer(算子实现层,第 2 层内部)
 ├── 接口层:参数校验 + 格式转换
 ├── 调度优化层:Tiling + 多核分配
 └── 算子实现层:Ascend C Kernel → NPU 执行

ATB 负责"要不要融合"、“怎么调接口”;ops-transformer 负责"融合后怎么算"、“Tiling 怎么配”。

所以如果你发现 FlashAttention 跑了但没加速(没融合成功),问题大概率在 ATB 的融合决策层(编译期的图融合没生效);如果你发现结果算错了(数值不对),问题大概率在 ops-transformer 的 Tiling 参数(Ub 溢出)。


总结:一句话说就是

ops-transformer 不是"一个算子库",是三层子架构(接口层 → 调度优化层 → 算子实现层)串联起来的 Transformer 类算子集合。FlashAttention 在里头从 Python API 到 NPU 执行,经过接口转换、Tiling 决策、在线 Softmax、可选融合、Runtime 下发五个阶段。快的真正原因是中间结果全程留在 Ub,最后一步才回 HBM——但这要求 Tiling 参数必须跟 Ub 容量严格匹配,配错了数值会静默出错,不报异常。

Logo

鲲鹏昇腾开发者社区是面向全社会开放的“联接全球计算开发者,聚合华为+生态”的社区,内容涵盖鲲鹏、昇腾资源,帮助开发者快速获取所需的知识、经验、软件、工具、算力,支撑开发者易学、好用、成功,成为核心开发者。

更多推荐