- 算子库
- 人工智能
- 深度学习
- Ascend
【免费下载链接】ops-transformer
本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。
SparseFlashMlaMetadata 是 CANN ops-transformer 开源仓库中SparseFlashMla稀疏 MLA(Multi-head Latent Attention)算子的前置调度算子:它本身不执行任何 Attention 数值计算,而是在 AI CPU 上根据各 Batch 的 Q/KV 序列长度与 mask 模式,通过开销模型为每个 AI Core 计算 Attention 计算任务的起止范围,输出固定 1024 个 INT32 的分核元数据(metadata),供主算子SparseFlashMla直接消费。本文基于 算子 README 与 aclnnSparseFlashMlaMetadata 接口文档,完整讲解该算子的功能定位、参数体系、平台约束、两段式 aclnn 调用方式与 PyTorch 调用方式,并结合 AI CPU 内核源码 剖析其负载均衡分核的底层实现。读完本文,你将能够独立为 SWA / CSA / HCA 三类稀疏注意力场景正确配置并调用该算子,理解 metadata 每个字段的业务含义,以及如何解析输出验证分核结果。
功能定位与典型使用场景
SparseFlashMlaMetadata服务于稀疏 MLA 注意力计算的完整流水线,其定位是主算子SparseFlashMla的前置调度算子。文档明确说明:
该算子不建议单独使用,建议与
SparseFlashMla算子配合使用,形成完整的工作流。
它解决的核心工程问题是负载不均衡:稀疏注意力中每个 query token 实际需要 attend 的 KV 范围差异巨大(滑动窗口、压缩稀疏、重度压缩三种模式各有不同的有效范围),如果简单地按序列长度均分给各 AI Core,会出现部分 Core 空闲、部分 Core 过载的情况。该算子根据输入参数在 AI CPU 计算出每个 AI Core 应处理的 Attention 计算起止范围,从而最大化计算资源的利用率,避免各 Core 间负载不均衡。
算子覆盖三类稀疏场景(场景简称):
| 场景简称 | 全称 | 含义 |
|---|---|---|
| SWA | Sliding Window Attention | 滑动窗口注意力,q 只关注窗口内的 ori_kv |
| CSA | Compressed Sparse Attention | 压缩稀疏注意力,q 关注经过压缩的 cmp_kv(压缩倍率 1、2 或 4) |
| HCA | Heavily Compressed Attention | 重度压缩注意力,cmp_kv 相对压缩前长度的压缩倍率为 128 |
工作原理与元数据输出格式
核心计算流程
SparseFlashMlaMetadata是 AICPU 调度算子,不涉及数值计算。接口文档给出其核心流程(四步):
- 解析各 Batch 的 Q/KV 序列长度:根据
layout_q/layout_kv与cu_seqlens_*、seqused_*、max_seqlen_*等输入推导每个 Batch 中 q、ori_kv、cmp_kv 的实际有效序列长度; - 根据 mask 模式计算每个 S1G 块的有效 S2 范围:mask 模式决定每个 S1G 行(S1×G 方向的分块)能访问的 KV 区间,稀疏场景下还叠加
ori_topk/cmp_topk/*_topk_length计算有效范围; - 基于开销模型进行负载均衡分核:以分块开销为单位,按核数切分任务,尽量使各核开销均衡;
- 输出分核元数据:将 FA(Flash Attention 计算)与 FD(Flash Decode 归约)两个阶段的任务切分结果写入固定大小的 metadata tensor。
metadata 输出结构
输出metadata的 shape 固定为(1024,),数据类型 INT32,内部划分为两大区域。示例代码 test_aclnn_sparse_flash_mla_metadata.cpp 中定义了对应的布局常量:
constexpr uint32_t AIC_CORE_MAX_NUM = 36; // FA 区域最多 36 个 AICore constexpr uint32_t AIV_CORE_MAX_NUM = 72; // FD 区域最多 72 个 AIVCore constexpr uint32_t SMLA_METADATA_TOTAL_SIZE = 1024; constexpr uint32_t FA_METADATA_SIZE = 9; constexpr uint32_t FD_METADATA_SIZE = 8;- FA Metadata 区域:
AIC_CORE_NUM × 9个 INT32,每个 AICore 的 FA(Flash Attention 前向计算)阶段任务信息:
| 索引 | 含义 |
|---|---|
| 0 | core_enable,该核是否启用 |
| 1 | bn2_start,BN2 起始索引 |
| 2 | m_start,M(S1G)起始索引 |
| 3 | s2_start,S2 起始索引 |
| 4 | bn2_end,BN2 结束索引 |
| 5 | m_end,M 结束索引 |
| 6 | s2_end,S2 结束索引 |
| 7 | first_fd_data_workspace_idx,第一份 FD 归约数据的 workspace 偏移 |
| 8 | max_s2_block_num,单核上分配到的最多的 s2 block 数 |
- FD Metadata 区域:
AIV_CORE_NUM × 8个 INT32,每个 AIVCore 的 FD(Flash Decode 归约)任务信息,前 7 个字段含义如下(索引 7 保留):
| 索引 | 含义 |
|---|---|
| 0 | core_enable,该核是否启用 |
| 1 | bn2_idx,归约任务的 BN2 索引 |
| 2 | m_idx,归约任务的 M 索引 |
| 3 | workspace_idx,归约数据在 workspace 中的存放位置 |
| 4 | workspace_num,S2 核间切分份数 |
| 5 | m_start,M 轴起点 |
| 6 | m_num,M 轴行数 |
符号说明
接口文档定义了以下核心符号,贯穿所有参数与约束描述:
| 符号 | 含义 |
|---|---|
| B | Batch Size |
| N1/N2 | Query/KV 头数 |
| D | 每个注意力头的维度 |
| G | GQA 分组比,G = N1/N2 |
| S1/S2 | Query/KV 序列长度 |
| S1G | S1×G 方向的分块索引 |
| mBaseSize | M 轴基本块大小,等于 G |
| s2BaseSize | S2 轴基本块大小,固定为 512 |
参数说明
算子共包含 3 个属性必选参数、9 个可选输入、9 个可选属性及 1 个输出。下表完整继承自 README,并按类别组织。
必选属性
| 参数名 | 输入/输出/属性 | 描述 | 数据类型 | 数据格式 |
|---|---|---|---|---|
| num_heads_q | 属性 | 表示q的头数,支持 [1, 128] | INT | - |
| num_heads_kv | 属性 | 表示ori_kv和cmp_kv的头数,仅支持 1 | INT | - |
| head_dim | 属性 | 注意力头的维度,仅支持 512 | INT | - |
可选输入
| 参数名 | 输入/输出/属性 | 描述 | 数据类型 | 数据格式 |
|---|---|---|---|---|
| cu_seqlens_q | 可选输入 | 表示 TND 布局下不同 batch 中q的累积序列长度,shape 为 (B+1,) | INT32 | ND |
| cu_seqlens_ori_kv | 可选输入 | 表示 TND 布局下不同 batch 中ori_kv的累积序列长度,shape 为 (B+1,) | INT32 | ND |
| cu_seqlens_cmp_kv | 可选输入 | 表示 TND 布局下不同 batch 中cmp_kv的累积序列长度,shape 为 (B+1,) | INT32 | ND |
| seqused_q | 可选输入 | 表示不同 batch 中q实际参与计算的 token 数,shape 为 (B,) | INT32 | ND |
| seqused_ori_kv | 可选输入 | 表示不同 batch 中ori_kv实际参与计算的 token 数,shape 为 (B,) | INT32 | ND |
| seqused_cmp_kv | 可选输入 | 表示不同 batch 中cmp_kv实际参与计算的 token 数,shape 为 (B,) | INT32 | ND |
| cmp_residual_kv | 可选输入 | 表示压缩 KV 余数,用于恢复 cmp 侧 mask 使用的压缩前 KV 长度,shape 为 (B,) | INT32 | ND |
| ori_topk_length | 可选输入 | SWA 稀疏 ori_kv 场景表示不同 q token 对应的 ori_kv 部分关键稀疏 token 的个数,必须传入,shape 为 (B, S1, N2) 或 (T1, N2) | INT32 | ND |
| cmp_topk_length | 可选输入 | 表示不同 q token 对应的 cmp_kv 部分关键稀疏 token 的个数,shape 为 (B, S1, N2) 或 (T1, N2) | INT32 | ND |
可选属性
| 参数名 | 输入/输出/属性 | 描述 | 数据类型 | 数据格式 |
|---|---|---|---|---|
| batch_size | 可选属性 | 表示输入样本批量大小;传入 0 时表示由接口推导,默认值为 0 | INT | - |
| max_seqlen_q | 可选属性 | 表示所有 batch 中q的最大有效 token 数;传入 0 时表示由接口推导,默认值为 0 | INT | - |
| max_seqlen_ori_kv | 可选属性 | 表示所有 batch 中ori_kv的最大有效 token 数;传入 0 时表示由接口推导,默认值为 0 | INT | - |
| max_seqlen_cmp_kv | 可选属性 | 表示所有 batch 中cmp_kv的最大有效 token 数;传入 0 时表示由接口推导,默认值为 0 | INT | - |
| ori_topk | 可选属性 | 表示从ori_kv中筛选出的关键稀疏 token 个数;SWA 稀疏 ori_kv 场景为主算子ori_sparse_indices最后一维 K 且必须大于 0,其他场景默认值为 0 | INT | - |
| cmp_topk | 可选属性 | 表示从cmp_kv中筛选出的关键稀疏 token 个数,默认值为 0 | INT | - |
| cmp_ratio | 可选属性 | 表示cmp_kv相对于压缩前 KV 长度的压缩倍率,用于恢复 cmp 侧 mask 使用的压缩前 KV 长度;仅传入ori_kv时不参与压缩 KV 计算。支持 [1, 128],默认值为 1 | INT | - |
| ori_mask_mode | 可选属性 | 表示q和ori_kv计算的 mask 模式,默认值为 0。0: No Mask;3: RightDownCausal 模式;4: Band 模式 | INT | - |
| cmp_mask_mode | 可选属性 | 表示q和cmp_kv计算的 mask 模式,默认值为 0。0: No Mask;3: RightDownCausal 模式 | INT | - |
| ori_win_left | 可选属性 | 表示q和ori_kv计算中q对过去 token 计算的数量,支持 -1 或非负数,其中 -1 表示窗口不受限,默认值为 -1 | INT | - |
| ori_win_right | 可选属性 | 表示q和ori_kv计算中q对未来 token 计算的数量,支持 -1 或非负数,其中 -1 表示窗口不受限,默认值为 -1 | INT | - |
| layout_q | 可选属性 | 表示输入q的数据排布格式,支持 "BSND" 和 "TND",默认值为 "BSND" | STRING | - |
| layout_kv | 可选属性 | 表示输入ori_kv和cmp_kv的数据排布格式,支持 "BSND"、"TND" 和 "PA_BBND",默认值为 "BSND" | STRING | - |
| has_ori_kv | 可选属性 | 表示SparseFlashMla主算子是否传入ori_kv,默认值为 true | BOOL | - |
| has_cmp_kv | 可选属性 | 表示SparseFlashMla主算子是否传入cmp_kv,默认值为 true | BOOL | - |
输出
| 参数名 | 输入/输出/属性 | 描述 | 数据类型 | 数据格式 |
|---|---|---|---|---|
| metadata | 输出 | 表示SparseFlashMla主算子使用的任务切分结果,shape 固定为 (1024,) | INT32 | ND |
需要特别说明两个容易混淆的参数对:
ori_topkvsori_topk_length:ori_topk是标量属性,表示从 ori_kv 中筛选的关键稀疏 token 总数(等于主算子ori_sparse_indices最后一维 K);ori_topk_length是逐 q token、逐 KV head 的实际有效索引条目数(左对齐),取值应在 [0, K] 范围内,SWA 稀疏 ori_kv 场景必须传入。分核时只使用ori_topk_length生成任务切分(详见下文平台约束)。cmp_ratio与cmp_residual_kv:压缩 KV 通过cmp_len * cmp_ratio + residual恢复压缩前的 KV 长度,cmp_residual_kv必须满足cmp_residual_kv[i] < cmp_ratio。
产品支持情况与平台差异
产品支持情况
| 产品 | 是否支持 |
|---|---|
| Ascend 950PR/Ascend 950DT | √ |
| Atlas A3 训练系列产品/Atlas A3 推理系列产品 | √ |
| Atlas A2 训练系列产品/Atlas A2 推理系列产品 | √ |
| Atlas 200I/500 A2 推理系列产品 | × |
| Atlas 推理系列产品 | × |
| Atlas 训练系列产品 | × |
平台差异要点
不同平台的参数支持范围存在显著差异,配置时务必对照:
- Atlas A3 / Atlas A2 系列:
num_heads_q/num_heads_kv仅支持 1、2、4、8、16、32、64、128;不支持seqused_q、cmp_topk_length;SWA 稀疏 ori_kv 场景支持ori_topk_length、ori_topk大于 0 及ori_mask_mode为 0,ori_win_left和ori_win_right支持非负数;其他 SWA 场景ori_topk为 0、ori_mask_mode为 4、ori_win_left为 127、ori_win_right为 0;cmp_topk支持 [0, 8192],cmp_mask_mode仅支持 3,SWA 不传入 cmp_kv,cmp_ratio不参与计算;CSA 场景cmp_ratio支持 1、2 或 4,HCA 场景支持 128。 - Ascend 950PR/Ascend 950DT:
cmp_ratio在 SWA 场景传 1,CSA/HCA 场景支持 1 到 128。
约束说明
通用规格约束
- 该接口支持训练、推理场景下使用,支持 aclgraph 模式。
- 符号约定:B(Batch)表示输入样本批量大小,q、ori_kv、cmp_kv 为配套的 SparseFlashMla 算子的入参;S1 表示 layout_q=BSND 时 q shape 中 S 轴的大小,T1 表示 layout_q=TND 时 q shape 中 T 轴的大小,S2 表示 layout_kv=BSND 时 ori_kv shape 中 S 轴的大小,S3 表示 layout_kv=BSND 时 cmp_kv shape 中 S 轴的大小,N2 表示 ori_kv、cmp_kv shape 中 N 轴的大小。
cu_seqlens_q、cu_seqlens_ori_kv、cu_seqlens_cmp_kv的值为当前 Batch 与前序 Batch 有效 token 数的累加值,第一个元素固定为 0,后一个元素的值必须大于等于前一个元素的值。seqused_q、seqused_ori_kv、seqused_cmp_kv的值表示每个 Batch 中的有效 token 数。layout_q和layout_kv组合仅支持 "BSND"/"BSND"、"TND"/"TND"、"BSND"/"PA_BBND"、"TND"/"PA_BBND";非 PA_BBND 场景下layout_q和layout_kv必须一致。cmp_residual_kv需满足cmp_residual_kv[i] < cmp_ratio。- aclnn 接口默认采用确定性实现,相同输入多次调用结果一致。
序列长度与 Batch 取值规则
算子的 seqlen 与 batch 推导遵循"显式输入优先、属性兜底"的规则,接口文档对此有严格定义:
- Batch 取值:layout_q 为 BSND 时,优先通过
seqused_q的 shape 推导 batch,未传入则通过batch_size获取;layout_q 为 TND 时,优先通过seqused_q的 shape 推导 batch,未传入则通过cu_seqlens_q的 shape 推导。 - q Seqlen 取值:layout_q 为 BSND 时,优先取
seqused_q元素,未传入则用max_seqlen_q;layout_q 为 TND 时,优先取seqused_q元素,未传入则用cu_seqlens_q元素。 - ori_kv Seqlen 取值:layout_kv 为 BSND 时优先
seqused_ori_kv、兜底max_seqlen_ori_kv;TND 时优先seqused_ori_kv、兜底cu_seqlens_ori_kv;PA_BBND 时优先seqused_ori_kv、兜底ori_topk_length。 - cmp_kv Seqlen 取值:layout_kv 为 BSND 时优先
seqused_cmp_kv、兜底max_seqlen_cmp_kv;TND 时优先seqused_cmp_kv、兜底cu_seqlens_cmp_kv;PA_BBND 时优先seqused_cmp_kv、兜底cmp_topk_length。 - BSND 布局下
max_seqlen_*必须显式传入:layout_q=BSND 时max_seqlen_q必须传 S1;has_ori_kv 为 true 时max_seqlen_ori_kv必须传 S2;has_cmp_kv 为 true 时max_seqlen_cmp_kv必须传 S3。 - TND 布局下
cu_seqlens_*必须显式传入:layout_q=TND 时cu_seqlens_q必须传入;has_ori_kv 为 true 时cu_seqlens_ori_kv必须传入;has_cmp_kv 为 true 时cu_seqlens_cmp_kv必须传入。
稀疏场景下的有效 seqlen 计算
对于 Ascend 950PR/Ascend 950DT,稀疏有效序列长度的判定规则为:
has_ori_kv为 true 时,ori_topk大于 0 认为 ori_kv 部分是稀疏的,为 0 则认为非稀疏;has_cmp_kv同理。- ori 侧:
ori_topk不为 0 且ori_mask_mode为 0 时,ori_topk_length必须传入,取ori_mask_mode规则与ori_topk_length元素的最小值作为当前 q token 对应的 ori_kv 有效 seqlen;其他 ori_kv 稀疏场景取ori_mask_mode规则与ori_topk的最小值。 - cmp 侧:
cmp_topk不为 0 且cmp_mask_mode为 0 时,cmp_topk_length必须传入,取 mask 规则与cmp_topk_length元素的最小值;其他 cmp_kv 稀疏场景取 mask 规则与cmp_topk的最小值。 - PA_BBND 布局下,
ori_topk_length必传场景中seqused_ori_kv可选传入,其他场景必须传入;cmp 侧同理。
SWA 稀疏 ori_kv 场景(Atlas A2/A3 系列)
- 仅支持 SWA 模板:
has_ori_kv为 true、has_cmp_kv为 false、ori_topk大于 0、ori_mask_mode为 0,ori_win_left和ori_win_right为非负数,且必须传入ori_topk_length。 ori_topk应与配套主算子ori_sparse_indices最后一维 K 保持一致;ori_topk_length表示每个 q token 和 KV head 的左对齐有效索引条目数,取值应在 [0, K] 范围内;Metadata 仅使用ori_topk_length生成任务切分。- 配套主算子在 PA_BBND 场景仍要求传入
seqused_ori_kv。 cmp_topk在 CSA 场景支持 [1, 8192] 内的任意整数,SWA、HCA 场景传 0;cmp_ratio在 CSA 场景传 1、2 或 4,HCA 场景传 128。
调用方式与示例
aclnn 两段式 API
aclnn 接口遵循 CANN 算子库的两段式调用规范:必须先调用aclnnSparseFlashMlaMetadataGetWorkspaceSize获取 workspace 大小与执行器,再调用aclnnSparseFlashMlaMetadata执行实际计算。函数原型(见 aclnn_sparse_flash_mla_metadata.h):
aclnnStatus aclnnSparseFlashMlaMetadataGetWorkspaceSize( const aclTensor *cuSeqlensQOptional, const aclTensor *cuSeqlensOriKvOptional, const aclTensor *cuSeqlensCmpKvOptional, const aclTensor *sequsedQOptional, const aclTensor *sequsedOriKvOptional, const aclTensor *sequsedCmpKvOptional, const aclTensor *cmpResidualKvOptional, const aclTensor *oriTopkLengthOptional, const aclTensor *cmpTopkLengthOptional, int64_t numHeadsQ, int64_t numHeadsKv, int64_t headDim, int64_t batchSize, int64_t maxSeqlenQ, int64_t maxSeqlenOriKv, int64_t maxSeqlenCmpKv, int64_t oriTopk, int64_t cmpTopk, int64_t cmpRatio, int64_t oriMaskMode, int64_t cmpMaskMode, int64_t oriWinLeft, int64_t oriWinRight, const char *layoutQOptional, const char *layoutKvOptional, bool hasOriKv, bool hasCmpKv, const aclTensor *metaData, uint64_t *workspaceSize, aclOpExecutor **executor); aclnnStatus aclnnSparseFlashMlaMetadata( void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream);第一段接口完成入参校验并创建执行器,常见返回值包括:
| 返回值 | 错误码 | 触发场景 |
|---|---|---|
| ACLNN_ERR_INNER_CREATE_EXECUTOR | 561101 | 创建 aclOpExecutor 失败 |
| ACLNN_ERR_INNER_NULLPTR | 561103 | workspaceSize 或 executor 为空指针;可选输入连续化后为空;添加 AICPU 任务失败 |
| ACLNN_ERR_PARAM_INVALID | 161002 | 参数不合法:如 batchSize/maxSeqlenQ 为负数、numHeadsQ 不在 [1,128]、numHeadsKv 不为 1、headDim 不为 512、mask 模式不合法、窗口值小于 -1、cmpRatio 不在 [1,128]、layout/cuSeqlens/seqused/metaData 的 shape 或必选关系不符等 |
从 aclnn 接口实现 可以看到,第一段接口内部会读取当前平台的 Cube 核数(GetCubeCoreNum)与 Vector 核数(GetVectorCoreNum),并通过aclrtGetSysParamOpt(ACL_OPT_DETERMINISTIC, ...)查询确定性开关,把socVersion、aicCoreNum、aivCoreNum、isBatchConsistency一并下传给 AICPU 内核;所有可选输入均会先做l0op::Contiguous连续化处理,因此非连续 Tensor 也被支持。
C++ 调用示例(HCA 场景)
完整可编译示例位于 test_aclnn_sparse_flash_mla_metadata.cpp,其核心调用流程如下(BSND 布局、HCA 场景,cmp_ratio=128):
// 1. device/stream 初始化(固定写法) int32_t deviceId = 0; aclrtStream stream; Init(deviceId, &stream); // 2. 构造输入与输出 // num_heads_q=64, num_heads_kv=1, head_dim=512 // ori_topk=0, cmp_topk=0, cmp_ratio=128 // ori_mask_mode=4 (Band), cmp_mask_mode=3 (RightDownCausal) // ori_win_left=127, ori_win_right=0 // layout_q="BSND", layout_kv="BSND" // has_ori_kv=true, has_cmp_kv=true // batch_size=4, max_seqlen_q=1024, max_seqlen_ori_kv=1024, max_seqlen_cmp_kv=1024 // metadata shape = {1024} // 3. 第一段接口:获取 workspace 大小与执行器 uint64_t workspaceSize = 0; aclOpExecutor *executor = nullptr; ret = aclnnSparseFlashMlaMetadataGetWorkspaceSize( cuSeqlensQOptional, cuSeqlensOriKvOptional, cuSeqlensCmpKvOptional, sequsedQOptional, sequsedOriKvOptional, sequsedCmpKvOptional, cmpResidualKvOptional, oriTopkLengthOptional, cmpTopkLengthOptional, numHeadsQ, numHeadsKv, headDim, batchSize, maxSeqlenQ, maxSeqlenOriKv, maxSeqlenCmpKv, oriTopk, cmpTopk, cmpRatio, oriMaskMode, cmpMaskMode, oriWinLeft, oriWinRight, layoutQOptional, layoutKvOptional, hasOriKv, hasCmpKv, metadata.data, &workspaceSize, &executor); // 按需申请 workspace void *workspaceAddr = nullptr; if (workspaceSize > 0) { aclrtMalloc(&workspaceAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); } // 4. 第二段接口:执行实际计算 ret = aclnnSparseFlashMlaMetadata(workspaceAddr, workspaceSize, executor, stream); // 5. 同步等待任务执行结束 aclrtSynchronizeStream(stream); // 6. 将 metadata 拷贝回 Host 并解析 SmlaMetadata result {}; aclrtMemcpy(&result, sizeof(result), metadata.deviceAddr, sizeof(result), ACL_MEMCPY_DEVICE_TO_HOST);示例中的SmlaMetadata结构体与 36/72 的核数上限均来自 稀疏 MLA 内核公共头文件,与上文 metadata 布局一一对应。示例还演示了cmpResidualKvOptional的构造条件:当hasCmpKv && cmpRatio != 1 && cmpMaskMode == 3时必须创建(CSA、HCA 且需要恢复压缩前长度)。
PyTorch 调用示例(CSA 场景)
通过torch.ops.cann_ops_transformer.sparse_flash_mla_metadata可直接在 PyTorch 中生成主算子使用的 metadata,示例见 test_torch_sparse_flash_mla_metadata.py。该示例是一个 TND + PA_BBND 布局的 CSA 场景(cmp_ratio=4):
import torch import torch_npu import torchair import cann_ops_transformer metadata = torch.ops.cann_ops_transformer.sparse_flash_mla_metadata( cu_seqlens_q = torch.tensor([0, 10], dtype=torch.int32).npu(), # (B+1,),首元素固定为 0 cu_seqlens_ori_kv = None, cu_seqlens_cmp_kv = None, seqused_q = None, seqused_ori_kv = torch.tensor([8192], dtype=torch.int32).npu(), # (B,) seqused_cmp_kv = torch.tensor([64], dtype=torch.int32).npu(), # (B,) cmp_residual_kv = torch.tensor([1], dtype=torch.int32).npu(), # 满足 residual < cmp_ratio ori_topk_length = None, cmp_topk_length = None, num_heads_q = 128, num_heads_kv = 1, head_dim = 512, batch_size = 1, max_seqlen_q = 1, max_seqlen_ori_kv = 512, max_seqlen_cmp_kv = 32, ori_topk = 0, cmp_topk = 512, # CSA 场景支持 [1, 8192] cmp_ratio = 4, # CSA 场景传 1、2 或 4 ori_mask_mode = 4, # Band cmp_mask_mode = 3, # RightDownCausal ori_win_left = 127, ori_win_right = 0, layout_q = "TND", layout_kv = "PA_BBND", has_ori_kv = True, has_cmp_kv = True )该接口生成的 metadata 会直接作为SparseFlashMla主算子的输入使用,完整的主算子 PyTorch 接口说明可参考 torchapi_sparse_flash_mla.md。
输出解析:验证分核结果
示例代码将 metadata 回拷 Host 后,按faMetadata[36][9]与fdMetadata[72][8]两个二维数组组织打印。对每个 AIC Core 打印Core Enable / Start BN2 / Start M / Start S2 / End BN2 / End M / End S2 / First Workspace Index / Max S2 Block Num;对每个 AIV Core 打印Core Enable / FD Task BN2 Idx / FD Task M Idx / FD Task Workspace Idx / FD Task Workspace Num / FD Subtask M Start / FD Subtask M Num。通过core_enable字段可以确认实际启用的核数,通过起止索引可以核对每个核被分配到的 (BN2, M, S2) 任务区间,从而验证负载均衡切分是否符合预期。
源码实现:AICPU 负载均衡调度原理
算子的核心逻辑位于 sparse_flash_mla_metadata_aicpu.h(及对应的.cpp实现)。从类SparseFlashMlaMetadataCpuKernel的成员方法可以完整还原其内部流水线:
- 准备阶段:
Prepare→ParamsCheck→ParamsInit,确定groupSize_、mBaseSize_、s2BaseSize_(默认 128,运行时按 block 切分动态推导)等内部属性,并分别调用CalcOriMaskMode/CalcCmpMaskMode归一化 mask 模式。 - 分块与开销计算:
CalcSplitInfo计算每个 batch 在 S1G、oriS2、cmpS2 三个方向切分出的基本块数与尾块 size;CalcBatchCost/CalcCostInfo汇总每个 batch 的总开销、总块数与最后一 block 的开销,维护totalBlockNum、totalCost、maxS1GCost等全局量。 - 开销模型:代码中定义了
COST_WEIGHT_M = 6与COST_WEIGHT_S2 = 10两个开销权重常量,以及FA_TOLERANCE_RATIO = 2的负载容差系数;block 开销按BlockType枚举(ORI_NORMAL_BLOCK/ORI_TAIL_BLOCK/CMP_NORMAL_BLOCK/CMP_TAIL_BLOCK)分类计算,通过BlockCost二维数组维护,体现了"普通块与尾块、ori 与 cmp 成本不同"的精细建模。 - 负载均衡分配:
BalanceSchedule驱动CalcSplitPlan,内部依次尝试AssignByBatch(按 batch 粗分)→AssignByRow(按 S1G 行分配)→AssignByBlock(按块细粒度分配)→ForceAssign(兜底强制分配),分配过程中用CoreCache(costLimit / cost / block / s2Loop)跟踪每个核的实时负载,最终由AssignBlocksToCore汇总为SplitResult。 - FD 归约任务切分:
SplitFD依据IsNeedRecordFDInfo/IsFirstReductionBlock识别需要进行核间归约的 block,调用RecordFDInfo记录归约任务的 BN2/M 索引、workspace 位置与 S2 切分份数,产出FlashDecodeResult(fdUsedVecNum、fdBN2Idx、fdMIdx、fdWorkspaceIdx、fdS2SplitNum、fdMSize 及每个 vector 的 fdMStart/fdMNum)。 - 元数据生成:
GenMetadata将SplitResult中的 FA 信息(usedCoreNum、bN2End、gS1End、s2End、firstFdDataWorkspaceIdx、maxCost、maxS2SplitNum 等)与 FD 信息按前述 9/8 字段布局写入metadata_输出 tensor。
可以看到,稀疏判定在源码层面对应isSparseOriKv_/isSparseCmpKv_/hasOriTopkLength_等状态位,mask 模式通过SparseMode枚举(DEFAULT_MASK / ALL_MASK / LEFT_UP_CAUSAL / RIGHT_DOWN_CAUSAL / BAND)表达,与文档中的ori_mask_mode/cmp_mask_mode取值一一对应。这正是 README 所述"根据输入参数在 AI CPU 计算出每个 AI Core 应处理的 Attention 计算起止范围"的代码级证据。
相关文档与代码索引
- 算子总览:attention/sparse_flash_mla_metadata/README.md
- aclnn 接口文档:attention/sparse_flash_mla_metadata/docs/aclnnSparseFlashMlaMetadata.md
- C++ 调用示例:attention/sparse_flash_mla_metadata/examples/test_aclnn_sparse_flash_mla_metadata.cpp
- PyTorch 调用示例:attention/sparse_flash_mla_metadata/examples/test_torch_sparse_flash_mla_metadata.py
- aclnn 接口声明与实现:aclnn_sparse_flash_mla_metadata.h、aclnn_sparse_flash_mla_metadata.cpp
- AICPU 内核实现:sparse_flash_mla_metadata_aicpu.h、sparse_flash_mla_metadata_aicpu.cpp
- 配套主算子元数据布局:attention/sparse_flash_mla/op_kernel/sparse_flash_mla_kernel_metadata.h
- 配套主算子 PyTorch 接口:attention/sparse_flash_mla/docs/torchapi_sparse_flash_mla.md
实际开发中,建议按"先确定平台(Ascend 950 / A2 / A3)→ 确定场景(SWA / CSA / HCA)→ 确定布局(BSND / TND / PA_BBND)→ 依据约束表逐项核对必传输入 → 两段式调用并解析 metadata"的顺序推进,即可把该调度算子稳定地接入稀疏 MLA 推理与训练链路。
- 算子库
- 人工智能
- 深度学习
- Ascend
【免费下载链接】ops-transformer
本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。
相关推荐
CANN ops-transformer 中 QuantSparseFlashMla 全量化稀疏 MLA 注意力算子实战指南
CANN ops transformer 中 QuantSparseFlashMla 全量化稀疏 MLA 注意力算子实战指南 本指南以 attention/qu
算子库人工智能深度学习AscendCANN ops-transformer QuantSparseFlashMlaMetadata 算子详解:量化稀疏 MLA 的 AI CPU 负载均衡分核方案
CANN ops transformer QuantSparseFlashMlaMetadata 算子详解:量化稀疏 MLA 的 AI CPU 负载均衡分核方案
算子库人工智能深度学习AscendCANN ops-transformer 稀疏注意力梯度元数据算子 SparseLightningIndexerKLLossGradMetadata 实战指南
CANN ops transformer 稀疏注意力梯度元数据算子 SparseLightningIndexerKLLossGradMetadata 实战指南
算子库人工智能深度学习Ascend
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考