CANN ops-transformer SparseFlashMlaMetadata 算子实战:稀疏 MLA 注意力负载均衡分核元数据生成指南
2026/9/21 15:16:11 网站建设 项目流程
  • 算子库
  • 人工智能
  • 深度学习
  • Ascend

【免费下载链接】ops-transformer

本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。

项目地址:https://gitcode.com/cann/ops-transformer
点击查看免费下载

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 间负载不均衡。

算子覆盖三类稀疏场景(场景简称):

场景简称全称含义
SWASliding Window Attention滑动窗口注意力,q 只关注窗口内的 ori_kv
CSACompressed Sparse Attention压缩稀疏注意力,q 关注经过压缩的 cmp_kv(压缩倍率 1、2 或 4)
HCAHeavily Compressed Attention重度压缩注意力,cmp_kv 相对压缩前长度的压缩倍率为 128

工作原理与元数据输出格式

核心计算流程

SparseFlashMlaMetadata是 AICPU 调度算子,不涉及数值计算。接口文档给出其核心流程(四步):

  1. 解析各 Batch 的 Q/KV 序列长度:根据layout_q/layout_kvcu_seqlens_*seqused_*max_seqlen_*等输入推导每个 Batch 中 q、ori_kv、cmp_kv 的实际有效序列长度;
  2. 根据 mask 模式计算每个 S1G 块的有效 S2 范围:mask 模式决定每个 S1G 行(S1×G 方向的分块)能访问的 KV 区间,稀疏场景下还叠加ori_topk/cmp_topk/*_topk_length计算有效范围;
  3. 基于开销模型进行负载均衡分核:以分块开销为单位,按核数切分任务,尽量使各核开销均衡;
  4. 输出分核元数据:将 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 前向计算)阶段任务信息:
索引含义
0core_enable,该核是否启用
1bn2_start,BN2 起始索引
2m_start,M(S1G)起始索引
3s2_start,S2 起始索引
4bn2_end,BN2 结束索引
5m_end,M 结束索引
6s2_end,S2 结束索引
7first_fd_data_workspace_idx,第一份 FD 归约数据的 workspace 偏移
8max_s2_block_num,单核上分配到的最多的 s2 block 数
  • FD Metadata 区域AIV_CORE_NUM × 8个 INT32,每个 AIVCore 的 FD(Flash Decode 归约)任务信息,前 7 个字段含义如下(索引 7 保留):
索引含义
0core_enable,该核是否启用
1bn2_idx,归约任务的 BN2 索引
2m_idx,归约任务的 M 索引
3workspace_idx,归约数据在 workspace 中的存放位置
4workspace_num,S2 核间切分份数
5m_start,M 轴起点
6m_num,M 轴行数

符号说明

接口文档定义了以下核心符号,贯穿所有参数与约束描述:

符号含义
BBatch Size
N1/N2Query/KV 头数
D每个注意力头的维度
GGQA 分组比,G = N1/N2
S1/S2Query/KV 序列长度
S1GS1×G 方向的分块索引
mBaseSizeM 轴基本块大小,等于 G
s2BaseSizeS2 轴基本块大小,固定为 512

参数说明

算子共包含 3 个属性必选参数、9 个可选输入、9 个可选属性及 1 个输出。下表完整继承自 README,并按类别组织。

必选属性

参数名输入/输出/属性描述数据类型数据格式
num_heads_q属性表示q的头数,支持 [1, 128]INT-
num_heads_kv属性表示ori_kvcmp_kv的头数,仅支持 1INT-
head_dim属性注意力头的维度,仅支持 512INT-

可选输入

参数名输入/输出/属性描述数据类型数据格式
cu_seqlens_q可选输入表示 TND 布局下不同 batch 中q的累积序列长度,shape 为 (B+1,)INT32ND
cu_seqlens_ori_kv可选输入表示 TND 布局下不同 batch 中ori_kv的累积序列长度,shape 为 (B+1,)INT32ND
cu_seqlens_cmp_kv可选输入表示 TND 布局下不同 batch 中cmp_kv的累积序列长度,shape 为 (B+1,)INT32ND
seqused_q可选输入表示不同 batch 中q实际参与计算的 token 数,shape 为 (B,)INT32ND
seqused_ori_kv可选输入表示不同 batch 中ori_kv实际参与计算的 token 数,shape 为 (B,)INT32ND
seqused_cmp_kv可选输入表示不同 batch 中cmp_kv实际参与计算的 token 数,shape 为 (B,)INT32ND
cmp_residual_kv可选输入表示压缩 KV 余数,用于恢复 cmp 侧 mask 使用的压缩前 KV 长度,shape 为 (B,)INT32ND
ori_topk_length可选输入SWA 稀疏 ori_kv 场景表示不同 q token 对应的 ori_kv 部分关键稀疏 token 的个数,必须传入,shape 为 (B, S1, N2) 或 (T1, N2)INT32ND
cmp_topk_length可选输入表示不同 q token 对应的 cmp_kv 部分关键稀疏 token 的个数,shape 为 (B, S1, N2) 或 (T1, N2)INT32ND

可选属性

参数名输入/输出/属性描述数据类型数据格式
batch_size可选属性表示输入样本批量大小;传入 0 时表示由接口推导,默认值为 0INT-
max_seqlen_q可选属性表示所有 batch 中q的最大有效 token 数;传入 0 时表示由接口推导,默认值为 0INT-
max_seqlen_ori_kv可选属性表示所有 batch 中ori_kv的最大有效 token 数;传入 0 时表示由接口推导,默认值为 0INT-
max_seqlen_cmp_kv可选属性表示所有 batch 中cmp_kv的最大有效 token 数;传入 0 时表示由接口推导,默认值为 0INT-
ori_topk可选属性表示从ori_kv中筛选出的关键稀疏 token 个数;SWA 稀疏 ori_kv 场景为主算子ori_sparse_indices最后一维 K 且必须大于 0,其他场景默认值为 0INT-
cmp_topk可选属性表示从cmp_kv中筛选出的关键稀疏 token 个数,默认值为 0INT-
cmp_ratio可选属性表示cmp_kv相对于压缩前 KV 长度的压缩倍率,用于恢复 cmp 侧 mask 使用的压缩前 KV 长度;仅传入ori_kv时不参与压缩 KV 计算。支持 [1, 128],默认值为 1INT-
ori_mask_mode可选属性表示qori_kv计算的 mask 模式,默认值为 0。0: No Mask;3: RightDownCausal 模式;4: Band 模式INT-
cmp_mask_mode可选属性表示qcmp_kv计算的 mask 模式,默认值为 0。0: No Mask;3: RightDownCausal 模式INT-
ori_win_left可选属性表示qori_kv计算中q对过去 token 计算的数量,支持 -1 或非负数,其中 -1 表示窗口不受限,默认值为 -1INT-
ori_win_right可选属性表示qori_kv计算中q对未来 token 计算的数量,支持 -1 或非负数,其中 -1 表示窗口不受限,默认值为 -1INT-
layout_q可选属性表示输入q的数据排布格式,支持 "BSND" 和 "TND",默认值为 "BSND"STRING-
layout_kv可选属性表示输入ori_kvcmp_kv的数据排布格式,支持 "BSND"、"TND" 和 "PA_BBND",默认值为 "BSND"STRING-
has_ori_kv可选属性表示SparseFlashMla主算子是否传入ori_kv,默认值为 trueBOOL-
has_cmp_kv可选属性表示SparseFlashMla主算子是否传入cmp_kv,默认值为 trueBOOL-

输出

参数名输入/输出/属性描述数据类型数据格式
metadata输出表示SparseFlashMla主算子使用的任务切分结果,shape 固定为 (1024,)INT32ND

需要特别说明两个容易混淆的参数对:

  • ori_topkvsori_topk_lengthori_topk是标量属性,表示从 ori_kv 中筛选的关键稀疏 token 总数(等于主算子ori_sparse_indices最后一维 K);ori_topk_length是逐 q token、逐 KV head 的实际有效索引条目数(左对齐),取值应在 [0, K] 范围内,SWA 稀疏 ori_kv 场景必须传入。分核时只使用ori_topk_length生成任务切分(详见下文平台约束)。
  • cmp_ratiocmp_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_qcmp_topk_length;SWA 稀疏 ori_kv 场景支持ori_topk_lengthori_topk大于 0 及ori_mask_mode为 0,ori_win_leftori_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 950DTcmp_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_qcu_seqlens_ori_kvcu_seqlens_cmp_kv的值为当前 Batch 与前序 Batch 有效 token 数的累加值,第一个元素固定为 0,后一个元素的值必须大于等于前一个元素的值。
  • seqused_qseqused_ori_kvseqused_cmp_kv的值表示每个 Batch 中的有效 token 数。
  • layout_qlayout_kv组合仅支持 "BSND"/"BSND"、"TND"/"TND"、"BSND"/"PA_BBND"、"TND"/"PA_BBND";非 PA_BBND 场景下layout_qlayout_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_leftori_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_EXECUTOR561101创建 aclOpExecutor 失败
ACLNN_ERR_INNER_NULLPTR561103workspaceSize 或 executor 为空指针;可选输入连续化后为空;添加 AICPU 任务失败
ACLNN_ERR_PARAM_INVALID161002参数不合法:如 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, ...)查询确定性开关,把socVersionaicCoreNumaivCoreNumisBatchConsistency一并下传给 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的成员方法可以完整还原其内部流水线:

  1. 准备阶段PrepareParamsCheckParamsInit,确定groupSize_mBaseSize_s2BaseSize_(默认 128,运行时按 block 切分动态推导)等内部属性,并分别调用CalcOriMaskMode/CalcCmpMaskMode归一化 mask 模式。
  2. 分块与开销计算CalcSplitInfo计算每个 batch 在 S1G、oriS2、cmpS2 三个方向切分出的基本块数与尾块 size;CalcBatchCost/CalcCostInfo汇总每个 batch 的总开销、总块数与最后一 block 的开销,维护totalBlockNumtotalCostmaxS1GCost等全局量。
  3. 开销模型:代码中定义了COST_WEIGHT_M = 6COST_WEIGHT_S2 = 10两个开销权重常量,以及FA_TOLERANCE_RATIO = 2的负载容差系数;block 开销按BlockType枚举(ORI_NORMAL_BLOCK/ORI_TAIL_BLOCK/CMP_NORMAL_BLOCK/CMP_TAIL_BLOCK)分类计算,通过BlockCost二维数组维护,体现了"普通块与尾块、ori 与 cmp 成本不同"的精细建模。
  4. 负载均衡分配BalanceSchedule驱动CalcSplitPlan,内部依次尝试AssignByBatch(按 batch 粗分)→AssignByRow(按 S1G 行分配)→AssignByBlock(按块细粒度分配)→ForceAssign(兜底强制分配),分配过程中用CoreCache(costLimit / cost / block / s2Loop)跟踪每个核的实时负载,最终由AssignBlocksToCore汇总为SplitResult
  5. FD 归约任务切分SplitFD依据IsNeedRecordFDInfo/IsFirstReductionBlock识别需要进行核间归约的 block,调用RecordFDInfo记录归约任务的 BN2/M 索引、workspace 位置与 S2 切分份数,产出FlashDecodeResult(fdUsedVecNum、fdBN2Idx、fdMIdx、fdWorkspaceIdx、fdS2SplitNum、fdMSize 及每个 vector 的 fdMStart/fdMNum)。
  6. 元数据生成GenMetadataSplitResult中的 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上加速计算。

项目地址:https://gitcode.com/cann/ops-transformer
点击查看免费下载
上一篇:猫抓视频嗅探工具:三步搞定网页视频下载的终极指南
下一篇:Rustup开发环境搭建:从源码编译到实战部署

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询