多彩编程 多彩编程MZPH · CODE BLOG
ARTICLE DETAIL

文章详情

深耕前端与后端开发技术的一线实战笔记与踩坑复盘。

CANN ops-transformer SparseFlashMlaMetadata 算子实战:稀疏 MLA 注意力负载均衡分核元数据生成指南

CANN ops-transformer SparseFlashMlaMetadata 算子实战:稀疏 MLA 注意力负载均衡分核元数据生成指南 算子库人工智能深度学习Ascend【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址https://gitcode.com/cann/ops-transformer点击查看免费下载SparseFlashMlaMetadata 是 CANN ops-transformer 开源仓库中SparseFlashMla稀疏 MLAMulti-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_kvCSACompressed Sparse Attention压缩稀疏注意力q 关注经过压缩的 cmp_kv压缩倍率 1、2 或 4HCAHeavily 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计算有效范围基于开销模型进行负载均衡分核以分块开销为单位按核数切分任务尽量使各核开销均衡输出分核元数据将 FAFlash Attention 计算与 FDFlash 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 的 FAFlash Attention 前向计算阶段任务信息索引含义0core_enable该核是否启用1bn2_startBN2 起始索引2m_startMS1G起始索引3s2_startS2 起始索引4bn2_endBN2 结束索引5m_endM 结束索引6s2_endS2 结束索引7first_fd_data_workspace_idx第一份 FD 归约数据的 workspace 偏移8max_s2_block_num单核上分配到的最多的 s2 block 数FD Metadata 区域AIV_CORE_NUM × 8个 INT32每个 AIVCore 的 FDFlash Decode 归约任务信息前 7 个字段含义如下索引 7 保留索引含义0core_enable该核是否启用1bn2_idx归约任务的 BN2 索引2m_idx归约任务的 M 索引3workspace_idx归约数据在 workspace 中的存放位置4workspace_numS2 核间切分份数5m_startM 轴起点6m_numM 轴行数符号说明接口文档定义了以下核心符号贯穿所有参数与约束描述符号含义BBatch SizeN1/N2Query/KV 头数D每个注意力头的维度GGQA 分组比G N1/N2S1/S2Query/KV 序列长度S1GS1×G 方向的分块索引mBaseSizeM 轴基本块大小等于 Gs2BaseSizeS2 轴基本块大小固定为 512参数说明算子共包含 3 个属性必选参数、9 个可选输入、9 个可选属性及 1 个输出。下表完整继承自 README并按类别组织。必选属性参数名输入/输出/属性描述数据类型数据格式num_heads_q属性表示q的头数支持 [1, 128]INT-num_heads_kv属性表示ori_kv和cmp_kv的头数仅支持 1INT-head_dim属性注意力头的维度仅支持 512INT-可选输入参数名输入/输出/属性描述数据类型数据格式cu_seqlens_q可选输入表示 TND 布局下不同 batch 中q的累积序列长度shape 为 (B1,)INT32NDcu_seqlens_ori_kv可选输入表示 TND 布局下不同 batch 中ori_kv的累积序列长度shape 为 (B1,)INT32NDcu_seqlens_cmp_kv可选输入表示 TND 布局下不同 batch 中cmp_kv的累积序列长度shape 为 (B1,)INT32NDseqused_q可选输入表示不同 batch 中q实际参与计算的 token 数shape 为 (B,)INT32NDseqused_ori_kv可选输入表示不同 batch 中ori_kv实际参与计算的 token 数shape 为 (B,)INT32NDseqused_cmp_kv可选输入表示不同 batch 中cmp_kv实际参与计算的 token 数shape 为 (B,)INT32NDcmp_residual_kv可选输入表示压缩 KV 余数用于恢复 cmp 侧 mask 使用的压缩前 KV 长度shape 为 (B,)INT32NDori_topk_length可选输入SWA 稀疏 ori_kv 场景表示不同 q token 对应的 ori_kv 部分关键稀疏 token 的个数必须传入shape 为 (B, S1, N2) 或 (T1, N2)INT32NDcmp_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可选属性表示q和ori_kv计算的 mask 模式默认值为 0。0: No Mask3: RightDownCausal 模式4: Band 模式INT-cmp_mask_mode可选属性表示q和cmp_kv计算的 mask 模式默认值为 0。0: No Mask3: RightDownCausal 模式INT-ori_win_left可选属性表示q和ori_kv计算中q对过去 token 计算的数量支持 -1 或非负数其中 -1 表示窗口不受限默认值为 -1INT-ori_win_right可选属性表示q和ori_kv计算中q对未来 token 计算的数量支持 -1 或非负数其中 -1 表示窗口不受限默认值为 -1INT-layout_q可选属性表示输入q的数据排布格式支持 BSND 和 TND默认值为 BSNDSTRING-layout_kv可选属性表示输入ori_kv和cmp_kv的数据排布格式支持 BSND、TND 和 PA_BBND默认值为 BSNDSTRING-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最后一维 Kori_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_lengthSWA 稀疏 ori_kv 场景支持ori_topk_length、ori_topk大于 0 及ori_mask_mode为 0ori_win_left和ori_win_right支持非负数其他 SWA 场景ori_topk为 0、ori_mask_mode为 4、ori_win_left为 127、ori_win_right为 0cmp_topk支持 [0, 8192]cmp_mask_mode仅支持 3SWA 不传入 cmp_kvcmp_ratio不参与计算CSA 场景cmp_ratio支持 1、2 或 4HCA 场景支持 128。Ascend 950PR/Ascend 950DTcmp_ratio在 SWA 场景传 1CSA/HCA 场景支持 1 到 128。约束说明通用规格约束该接口支持训练、推理场景下使用支持 aclgraph 模式。符号约定BBatch表示输入样本批量大小q、ori_kv、cmp_kv 为配套的 SparseFlashMla 算子的入参S1 表示 layout_qBSND 时 q shape 中 S 轴的大小T1 表示 layout_qTND 时 q shape 中 T 轴的大小S2 表示 layout_kvBSND 时 ori_kv shape 中 S 轴的大小S3 表示 layout_kvBSND 时 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_qlayout_q 为 TND 时优先取seqused_q元素未传入则用cu_seqlens_q元素。ori_kv Seqlen 取值layout_kv 为 BSND 时优先seqused_ori_kv、兜底max_seqlen_ori_kvTND 时优先seqused_ori_kv、兜底cu_seqlens_ori_kvPA_BBND 时优先seqused_ori_kv、兜底ori_topk_length。cmp_kv Seqlen 取值layout_kv 为 BSND 时优先seqused_cmp_kv、兜底max_seqlen_cmp_kvTND 时优先seqused_cmp_kv、兜底cu_seqlens_cmp_kvPA_BBND 时优先seqused_cmp_kv、兜底cmp_topk_length。BSND 布局下max_seqlen_*必须显式传入layout_qBSND 时max_seqlen_q必须传 S1has_ori_kv 为 true 时max_seqlen_ori_kv必须传 S2has_cmp_kv 为 true 时max_seqlen_cmp_kv必须传 S3。TND 布局下cu_seqlens_*必须显式传入layout_qTND 时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为 0ori_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 场景传 0cmp_ratio在 CSA 场景传 1、2 或 4HCA 场景传 128。调用方式与示例aclnn 两段式 APIaclnn 接口遵循 CANN 算子库的两段式调用规范必须先调用aclnnSparseFlashMlaMetadataGetWorkspaceSize获取 workspace 大小与执行器再调用aclnnSparseFlashMlaMetadata执行实际计算。函数原型见 aclnn_sparse_flash_mla_metadata.haclnnStatus 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, ...)查询确定性开关把socVersion、aicCoreNum、aivCoreNum、isBatchConsistency一并下传给 AICPU 内核所有可选输入均会先做l0op::Contiguous连续化处理因此非连续 Tensor 也被支持。C 调用示例HCA 场景完整可编译示例位于 test_aclnn_sparse_flash_mla_metadata.cpp其核心调用流程如下BSND 布局、HCA 场景cmp_ratio128// 1. device/stream 初始化固定写法 int32_t deviceId 0; aclrtStream stream; Init(deviceId, stream); // 2. 构造输入与输出 // num_heads_q64, num_heads_kv1, head_dim512 // ori_topk0, cmp_topk0, cmp_ratio128 // ori_mask_mode4 (Band), cmp_mask_mode3 (RightDownCausal) // ori_win_left127, ori_win_right0 // layout_qBSND, layout_kvBSND // has_ori_kvtrue, has_cmp_kvtrue // batch_size4, max_seqlen_q1024, max_seqlen_ori_kv1024, max_seqlen_cmp_kv1024 // 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_ratio4import 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], dtypetorch.int32).npu(), # (B1,)首元素固定为 0 cu_seqlens_ori_kv None, cu_seqlens_cmp_kv None, seqused_q None, seqused_ori_kv torch.tensor([8192], dtypetorch.int32).npu(), # (B,) seqused_cmp_kv torch.tensor([64], dtypetorch.int32).npu(), # (B,) cmp_residual_kv torch.tensor([1], dtypetorch.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 三个方向切分出的基本块数与尾块 sizeCalcBatchCost/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兜底强制分配分配过程中用CoreCachecostLimit / cost / block / s2Loop跟踪每个核的实时负载最终由AssignBlocksToCore汇总为SplitResult。FD 归约任务切分SplitFD依据IsNeedRecordFDInfo/IsFirstReductionBlock识别需要进行核间归约的 block调用RecordFDInfo记录归约任务的 BN2/M 索引、workspace 位置与 S2 切分份数产出FlashDecodeResultfdUsedVecNum、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.mdaclnn 接口文档attention/sparse_flash_mla_metadata/docs/aclnnSparseFlashMlaMetadata.mdC 调用示例attention/sparse_flash_mla_metadata/examples/test_aclnn_sparse_flash_mla_metadata.cppPyTorch 调用示例attention/sparse_flash_mla_metadata/examples/test_torch_sparse_flash_mla_metadata.pyaclnn 接口声明与实现aclnn_sparse_flash_mla_metadata.h、aclnn_sparse_flash_mla_metadata.cppAICPU 内核实现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点击查看免费下载相关推荐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上一篇猫抓视频嗅探工具三步搞定网页视频下载的终极指南下一篇Rustup开发环境搭建从源码编译到实战部署创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表