ARTICLE DETAIL

资讯详情

深耕编程入门与网站建设的一线实战洞察。

MixedQuantSparseFlashMlaMetadata 算子解析:CANN ops-transformer 中混合量化稀疏 MLA 的 AI CPU 负载均衡前置算子

MixedQuantSparseFlashMlaMetadata 算子解析:CANN ops-transformer 中混合量化稀疏 MLA 的 AI CPU 负载均衡前置算子 MixedQuantSparseFlashMlaMetadata 算子解析CANN ops-transformer 中混合量化稀疏 MLA 的 AI CPU 负载均衡前置算子【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformerMixedQuantSparseFlashMlaMetadata是 CANN ops-transformer 算子库中MixedQuantSparseFlashMla算子的前置Metadata 生成算子。它本身不执行任何 Attention 计算而是在 AI CPU 上根据输入的形状、Sequence Length、mask 模式与稀疏 TopK 信息为后续真正的 Attention 算子计算出每个 AI CoreAIC应处理的起止范围以及每个 AIV 核上 FDFlash Decode归约任务的分配方案从而避免各 Core 间负载不均。读完本文你将掌握该算子的功能定位、全部输入输出与属性参数的取值范围、配套使用约束并能够通过 aclnn API 或 PyTorch API 正确调用它同时理解其负载均衡算法的源码级实现原理。功能定位Attention 计算之前的“调度器”在长序列、大 Batch 的 MLAMulti-head Latent Attention推理与训练场景中不同 Batch、不同 Query token 对应的有效 KV 长度差异巨大若将计算任务简单均分到多个 AI Core必然出现严重的负载不均衡。MixedQuantSparseFlashMlaMetadata正是为了解决这一问题而存在它是MixedQuantSparseFlashMla算子的前置算子二者必须配套使用见 约束说明。它接受MixedQuantSparseFlashMla算子输入数据的 shape 信息batchSize、q Seqlen、kv Seqlen、mask 等在 AI CPU 上通过对输入分块并模拟计算耗时将分块均匀分配到可用的 AIC/AIV 核上。分配结果写入一个 shape 固定为(1024,)的 INT32 输出张量metadata后续作为MixedQuantSparseFlashMla算子的输入使用。从源码结构看该算子的计算完全在 AI CPU 侧完成kernel 实现位于 op_kernel_aicpu/mixed_quant_sparse_flash_mla_metadata_aicpu.cpp其主流程Compute()依次执行参数准备Prepare、负载均衡调度BalanceSchedule与 metadata 生成GenMetadata三步最终通过REGISTER_CPU_KERNEL注册为名为MixedQuantSparseFlashMlaMetadata的 CPU kernel。官方文档明确提示该算子不建议单独使用建议与aclnnMixedQuantSparseFlashMla算子配合使用形成完整的工作流。产品支持情况产品是否支持Ascend 950PR/Ascend 950DT√Atlas A3 训练系列产品/Atlas A3 推理系列产品×Atlas A2 训练系列产品/Atlas A2 推理系列产品×Atlas 200I/500 A2 推理产品×Atlas 推理系列产品×Atlas 训练系列产品×当前仅支持 Ascend 950 系列950PR/950DT。kernel 源码中的ProcessSocVersion()会检测 soc_version 中是否包含Ascend950前缀据此区分ASCEND950与ASCEND910两种平台分支950 平台在groupSize 64时还会将 AIC 核数减半以平衡负载详见 源码级原理。参数说明算子的输入/输出/属性参数如下表。其中 BBatch表示输入样本批量大小q、ori_kv、cmp_kv 为配套的MixedQuantSparseFlashMla算子的入参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可选输入表示不同 Batch 中 q 的有效 Sequence Lengthshape 为 (B1, )。INT32NDcu_seqlens_ori_kv可选输入表示不同 Batch 中 ori_kv 的有效 Sequence Lengthshape 为 (B1, )。INT32NDcu_seqlens_cmp_kv可选输入表示不同 Batch 中 cmp_kv 的有效 Sequence Lengthshape 为 (B1, )。INT32NDseqused_q可选输入表示不同 Batch 中 q 实际参与运算的 Sequence Lengthshape 为 (B, )。INT32NDseqused_ori_kv可选输入表示不同 Batch 中 ori_kv 实际参与运算的 Sequence Lengthshape 为 (B, )。INT32NDseqused_cmp_kv可选输入表示不同 Batch 中 cmp_kv 实际参与运算的 Sequence Lengthshape 为 (B, )。INT32NDcmp_residual_kv可选输入表示不同 Batch 中 cmp_kv 压缩后 Sequence Length 的余数配合 cmp_ratio 实现 cmp_kv 部分的 mask 和负载计算。cmp_mask_mode3 且 cmp_ratio≠1 时必须传入shape 为 (B, )。INT32NDori_topk_length可选输入表示不同 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)。INT32NDnum_heads_q属性表示 q 的 head 个数当前支持 [1, 128]。INT32-num_heads_kv属性表示 ori_kv、cmp_kv 对应的多头数当前仅支持 1。INT32-head_dim属性表示注意力头的维度当前仅支持 512。INT32-quant_mode属性表示量化模式1 表示 K、V nope 为 per-token-group 量化scale 类型为 bfloat162 表示 K、V nope 为 per-token-group 量化scale 类型为 float8_e8m0。INT32-batch_size可选属性表示 Batch 数量默认值为 0。INT32-max_seqlen_q可选属性表示 q 的最长 Sequence Length默认值为 0。INT32-max_seqlen_ori_kv可选属性表示 ori_kv 的最长 Sequence Length默认值为 0。INT32-max_seqlen_cmp_kv可选属性表示 cmp_kv 的最长 Sequence Length默认值为 0。INT32-ori_topk可选属性表示 ori_kv 中筛选出的关键稀疏 token 的个数0 表示非稀疏场景默认值为 0。INT32-cmp_topk可选属性表示 cmp_kv 中筛选出的关键稀疏 token 的个数0 表示非稀疏场景默认值为 0。INT32-rope_head_dim可选属性表示 rope 头的维度默认值为 64当前仅支持 64。INT32-cmp_ratio可选属性表示对 cmp_kv 的压缩率默认值为 1当前支持 [1, 128]。INT32-ori_mask_mode可选属性表示 q 和 ori_kv 计算的 mask 模式当前支持 0、3、4。0 表示 No Mask3 表示 RightDownCausal 模式4 表示 sliding window 模式默认值为 0。INT32-cmp_mask_mode可选属性表示 q 和 cmp_kv 计算的 mask 模式当前支持 0、3。0 表示 No Mask3 表示 RightDownCausal 模式默认值为 0。INT32-ori_win_left可选属性表示 q 和 ori_kv 计算中 q 对过去 token 计算的数量-1 表示无穷大默认值为 -1。INT32-ori_win_right可选属性表示 q 和 ori_kv 计算中 q 对未来 token 计算的数量-1 表示无穷大默认值为 -1。INT32-layout_q可选属性表示 q 的排列格式支持 BSND、TND默认值为 BSND。STRING-layout_kv可选属性表示 ori_kv、cmp_kv 的排列格式支持 BSND、TND、PA_BBND默认值为 BSND。STRING-has_ori_kv可选属性用于标识是否含有 ori_kv默认值为 true。BOOL-has_cmp_kv可选属性用于标识是否含有 cmp_kv默认值为 true。BOOL-metadata输出表示负载均衡结果输出shape 固定为 (1024, )。INT32ND关键属性取值语义补充quant_mode 两种量化布局来自 aclnnMixedQuantSparseFlashMlaMetadata.mdquant_mode1Q 的 noperope 非量化KV 的 nope 为 per-token-group FP8_e4m3 量化group_size64rope 部分非量化且与 q 保持一致scale 类型为 bf16kv_cache_layout 为block_size * (rope[64*2] nope[448] scale[448/64*2] pad[18])。quant_mode2Q 的 noperope 非量化KV 的 nope 为 per-token-group FP8_e4m3 量化group_size64rope 部分非量化且与 q 保持一致scale 类型为 e8m0kv_cache_layout 为block_size * (nope rope) block_size * (scale pad[1])。mask 模式详情ori_mask_mode 与 cmp_mask_mode 所表示的 mask 模式No Mask、RightDownCausal、sliding window 等的详细介绍见 sparse_mode 参数说明。在 host 侧参数校验源码 mixed_quant_sparse_flash_mla_metadata_check.h 中num_heads_q的合法区间被定义为[1, 128]、num_heads_kv必须等于 1、head_dim必须等于 512、cmp_ratio的合法区间为[1, 128]、quant_mode的合法区间为[1, 2]这些与文档中的约束一致同时该文件要求num_heads_q必须能被num_heads_kv整除。约束说明以下约束与取值规则必须严格遵守否则算子校验失败TND 场景下 AI CPU kernel 的ParamsCheck()会逐一校验 cu_seqlens 首元素为 0 且单调递增、seqused 各元素非负且不超过对应 seqlen 等详见 mixed_quant_sparse_flash_mla_metadata_aicpu.cppMixedQuantSparseFlashMlaMetadata 算子需要与 MixedQuantSparseFlashMla 算子配套使用。BBatch表示输入样本批量大小q、ori_kv、cmp_kv 为配套的 MixedQuantSparseFlashMla 算子的入参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 数。参数 cmp_residual_kv 需满足 cmp_residual_kv[i] cmp_ratio。ori_mask_mode 及 cmp_mask_mode 所表示的 mask 模式的详细介绍见 sparse_mode 参数说明。非 PA 场景 layout_q、layout_kv 须相同。has_ori_kv 为 true 时ori_topk 大于 0 认为 ori_kv 部分是稀疏的ori_topk 为 0 则认为 ori_kv 部分是非稀疏的。has_cmp_kv 为 true 时cmp_topk 大于 0 认为 cmp_kv 部分是稀疏的cmp_topk 为 0 则认为 cmp_kv 部分是非稀疏的。has_ori_kv 为 trueori_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 的最小值作为当前 q token 对应的 ori_kv 的有效 seqlen。has_cmp_kv 为 truecmp_topk 不为 0 且 cmp_mask_mode 为 0 时cmp_topk_length 必须传入此时取 cmp_mask_mode 规则与 cmp_topk_length 元素的最小值作为当前 q token 对应的 cmp_kv 的有效 seqlen其他 cmp_kv 稀疏场景取 cmp_mask_mode 规则与 cmp_topk 的最小值作为当前 q token 对应的 cmp_kv 的有效 seqlen。layout_qBSND 场景max_seqlen_q 必须传入 S1 的值。layout_kvBSND 场景has_ori_kv 为 true 时max_seqlen_ori_kv 必须传入 S2 的值has_cmp_kv 为 true 时max_seqlen_cmp_kv 必须传入 S3 的值。layout_qTND 场景cu_seqlens_q 必须传入。layout_kvTND 场景has_ori_kv 为 true 时cu_seqlens_ori_kv 必须传入has_cmp_kv 为 true 时cu_seqlens_cmp_kv 必须传入。layout_kvPA_BBND 场景has_ori_kv 为 trueori_topk 不为 0 且 ori_mask_mode 为 0 时ori_topk_length 必传场景seqused_ori_kv 可选传入其他场景 seqused_ori_kv 必须传入。has_cmp_kv 为 truecmp_topk 不为 0 且 cmp_mask_mode 为 0 时cmp_topk_length 必传场景seqused_cmp_kv 可选传入其他场景 seqused_cmp_kv 必须传入。Batch 取值规则layout_q 为 BSND 时优先通过 seqused_q 的 shape 推导 batchseqused_q 未传入则通过 batch_size 获取 batch 数。layout_q 为 TND 时优先通过 seqused_q 的 shape 推导 batchseqused_q 未传入则通过 cu_seqlens_q 的 shape 推导 batch。q Seqlen 取值规则layout_q 为 BSND 时优先通过 seqused_q 中的元素获取 seqlenseqused_q 未传入则通过 max_seqlen_q 获取 seqlen。layout_q 为 TND 时优先通过 seqused_q 中的元素获取 seqlenseqused_q 未传入则通过 cu_seqlens_q 中的元素获取 seqlen。ori_kv Seqlen 取值规则layout_kv 为 BSND 时优先通过 seqused_ori_kv 中的元素获取 seqlenseqused_ori_kv 未传入则通过 max_seqlen_ori_kv 获取 seqlen。layout_kv 为 TND 时优先通过 seqused_ori_kv 中的元素获取 seqlenseqused_ori_kv 未传入则通过 cu_seqlens_ori_kv 中的元素获取 seqlen。layout_kv 为 PA_BBND 时优先通过 seqused_ori_kv 中的元素获取 seqlenseqused_ori_kv 未传入则通过 ori_topk_length 获取 seqlen。cmp_kv Seqlen 取值规则layout_kv 为 BSND 时优先通过 seqused_cmp_kv 中的元素获取 seqlenseqused_cmp_kv 未传入则通过 max_seqlen_cmp_kv 获取 seqlen。layout_kv 为 TND 时优先通过 seqused_cmp_kv 中的元素获取 seqlenseqused_cmp_kv 未传入则通过 cu_seqlens_cmp_kv 中的元素获取 seqlen。layout_kv 为 PA_BBND 时优先通过 seqused_cmp_kv 中的元素获取 seqlenseqused_cmp_kv 未传入则通过 cmp_topk_length 获取 seqlen。上述 seqlen 的“多级回退”取值逻辑在 kernel 源码中均有对应实现例如GetS1SeqSize()依次尝试seqused_q→ TND 下的cu_seqlens_q→max_seqlen_qGetOriS2SeqSize()与GetCmpS2SeqSize()在 PA 场景或稀疏场景下还会回退到 UINT32_MAX表示需按 TopK 计算GetQueryBatchSize()则实现 batch 的推导顺序。调用方式该算子提供两套调用入口CANN 的 aclnn API 与 PyTorch API。仓库中提供了完整可运行的示例test_aclnn_mixed_quant_sparse_flash_mla_metadata.cpp 与 test_torch_mixed_quant_sparse_flash_mla_metadata.py。aclnn API两段式调用与其他 aclnn 算子一致aclnnMixedQuantSparseFlashMlaMetadata采用两段式接口先调用aclnnMixedQuantSparseFlashMlaMetadataGetWorkspaceSize获取 workspace 大小再调用aclnnMixedQuantSparseFlashMlaMetadata执行计算。aclnnStatus aclnnMixedQuantSparseFlashMlaMetadataGetWorkspaceSize( 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 quantMode, int64_t batchSize, int64_t maxSeqlenQ, int64_t maxSeqlenOriKv, int64_t maxSeqlenCmpKv, int64_t oriTopk, int64_t cmpTopk, int64_t ropeHeadDim, 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 aclnnMixedQuantSparseFlashMlaMetadata( void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream)第一段接口的参数要点9 个可选输入cuSeqlensQOptional等均支持空 Tensor传nullptr即可传入时要求 ND 格式、INT32、1 维ori_topk_length/cmp_topk_length 为 2 维或 3 维并支持非连续 Tensor框架内部会先执行 Contiguous 处理对应 aclnn_mixed_quant_sparse_flash_mla_metadata.cpp 中l0op::Contiguous的调用序列。数值/字符串/布尔属性参数必须落在上一节列出的支持范围内否则第一段接口返回ACLNN_ERR_PARAM_INVALID错误码 161002。输出metadata为 ND 格式、INT32、shape 固定(1024,)不支持非连续 Tensor。workspaceSizeuint64_t*返回需要在 Device 侧申请的 workspace 大小executoraclOpExecutor**返回包含算子计算流程的执行器。第一段接口还可能返回ACLNN_ERR_INNER_CREATE_EXECUTOR561101创建 aclOpExecutor 失败与ACLNN_ERR_INNER_NULLPTR561103workspaceSize/executor 为空指针或 Contiguous 处理后输入为空指针。返回值与 aclnn 返回码的完整含义参见 aclnn 返回码说明。第二段接口的 4 个参数为workspaceDevice 侧 workspace 内存地址、workspaceSize由第一段接口获取、executorop 执行器、stream指定执行任务的 Stream。该算子为默认确定性实现。aclnn 调用示例关键步骤以下为官方示例 test_aclnn_mixed_quant_sparse_flash_mla_metadata.cpp 的关键流程完整可编译代码请直接查看示例文件编译与运行方法见编译与运行样例// 1. 初始化 device 与 stream固定写法 aclInit(nullptr); aclrtSetDevice(deviceId); aclrtCreateStream(stream); // 2. 构造输入输出 // numHeadsQ64, numHeadsKv1, headDim512, quantMode1 // oriTopk0, cmpTopk0, ropeHeadDim64, cmpRatio128 // oriMaskMode4(sliding window), cmpMaskMode3(RightDownCausal) // oriWinLeft127, oriWinRight0 // layoutQBSND, layoutKvBSND, hasOriKvtrue, hasCmpKvtrue // batchSize4, maxSeqlenQmaxSeqlenOriKvmaxSeqlenCmpKv1024 // metadata: INT32, shape {1024} // 可选输入传 nullptr若 hasCmpKv cmpRatio!1 cmpMaskMode3需传 cmpResidualKvshape {B} // 3. 第一段接口获取 workspace 大小与 executor aclnnMixedQuantSparseFlashMlaMetadataGetWorkspaceSize( cuSeqlensQOptional, cuSeqlensOriKvOptional, cuSeqlensCmpKvOptional, sequsedQOptional, sequsedOriKvOptional, sequsedCmpKvOptional, cmpResidualKvOptional, oriTopkLengthOptional, cmpTopkLengthOptional, numHeadsQ, numHeadsKv, headDim, quantMode, batchSize, maxSeqlenQ, maxSeqlenOriKv, maxSeqlenCmpKv, oriTopk, cmpTopk, ropeHeadDim, cmpRatio, oriMaskMode, cmpMaskMode, oriWinLeft, oriWinRight, layoutQOptional, layoutKvOptional, hasOriKv, hasCmpKv, metadata, workspaceSize, executor); // 按需在 Device 侧申请 workspace if (workspaceSize 0) { aclrtMalloc(workspaceAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); } // 4. 第二段接口执行计算 aclnnMixedQuantSparseFlashMlaMetadata(workspaceAddr, workspaceSize, executor, stream); // 5. 同步等待并取回结果 aclrtSynchronizeStream(stream); aclrtMemcpy(result, sizeof(result), metadata, sizeof(result), ACL_MEMCPY_DEVICE_TO_HOST);PyTorch API 调用示例通过torch.ops.cann_ops_transformer.mixed_quant_sparse_flash_mla_metadata(...)调用接口说明见 torchapi_mixed_quant_sparse_flash_mla.md。官方示例 test_torch_mixed_quant_sparse_flash_mla_metadata.py 展示了 TND PA_BBND 混合 layout、稀疏 cmp_kv 的典型用法import torch import torch_npu import torchair import cann_ops_transformer metadata torch.ops.cann_ops_transformer.mixed_quant_sparse_flash_mla_metadata( cu_seqlens_q torch.tensor([0, 10], dtypetorch.int32).npu(), # (B1,) 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(), # (B,)cmp_ratio!1 且 cmp_mask_mode3 时必传 ori_topk_length None, cmp_topk_length None, num_heads_q 128, num_heads_kv 1, head_dim 512, quant_mode 1, batch_size 1, max_seqlen_q 1, max_seqlen_ori_kv 512, max_seqlen_cmp_kv 32, ori_topk 0, # ori_kv 非稀疏 cmp_topk 512, # cmp_kv 稀疏关键 token 数 512 rope_head_dim 64, cmp_ratio 4, # cmp_kv 压缩率 4 ori_mask_mode 4, # sliding window 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 的结构解析输出张量metadata长度为 1024INT32内部按MqsmlaMetadata结构组织定义于 mixed_quant_sparse_flash_mla_metadata.h分为 FAFlash AttentionAIC 核任务与 FDFlash DecodeAIV 核归约任务两部分constexpr uint32_t AIC_CORE_MAX_NUM 36; constexpr uint32_t AIV_CORE_MAX_NUM 72; constexpr uint32_t MQSMLA_METADATA_TOTAL_SIZE 1024; constexpr uint32_t FA_METADATA_SIZE 9; constexpr uint32_t FD_METADATA_SIZE 8; struct MqsmlaMetadata { uint32_t faMetadata[AIC_CORE_MAX_NUM][FA_METADATA_SIZE]; // 每个 AIC 核 9 个字段 uint32_t fdMetadata[AIV_CORE_MAX_NUM][FD_METADATA_SIZE]; // 每个 AIV 核 8 个字段 };FA Metadata 每个 AIC 核的 9 个字段下标常量见示例文件下标字段含义0FA_CORE_ENABLE_INDEX该 AIC 核是否启用1 启用 / 0 停用1FA_BN2_START_INDEX起始 BN2Batch × N2索引2FA_M_START_INDEX起始 MS1G 行索引3FA_S2_START_INDEX起始 S2KV 块索引4FA_BN2_END_INDEX结束 BN2 索引5FA_M_END_INDEX结束 M 索引6FA_S2_END_INDEX结束 S2 索引7FA_FIRST_FD_DATA_WORKSPACE_IDX_INDEX首个 FD 数据 workspace 索引8FA_S2_MAX_NUM单核 S2 计算轮次最大值FD Metadata 每个 AIV 核的 8 个字段下标字段含义0FD_CORE_ENABLE_INDEX该 AIV 核是否启用1FD_BN2_IDX_INDEXFD 任务对应的 BN2 索引2FD_M_IDX_INDEXFD 任务对应的 M 索引3FD_WORKSPACE_IDX_INDEXFD 任务 workspace 起始索引4FD_WORKSPACE_NUM_INDEXFD 任务 workspace 数量5FD_M_START_INDEXFD 子任务 M 起始6FD_M_NUM_INDEXFD 子任务 M 数量在示例程序中取回结果后会按AIC_CORE_MAX_NUM36、AIV_CORE_MAX_NUM72两层循环分别打印每个 AIC 核的启停区间与每个 AIV 核的 FD 任务信息读者可据此验证调度结果。注意示例中的 36/72 为头文件约定的最大核数常量来自smla_metadata_common.h实际使用的核数由 AI CPU 侧根据平台信息aicCoreNum_/aivCoreNum_与负载量决定未使用的核其 Core Enable 置 0。源码级原理负载均衡调度如何工作AI CPU kernel 的实现mixed_quant_sparse_flash_mla_metadata_aicpu.cpp展示了完整的负载均衡算法可概括为“切分 → 计费 → 分配 → 写元数据”四个阶段划分基本块CalcSplitInfo按 batch 将 q 的 S1 轴以mBaseSize_默认等于 groupSize numHeadsQ/numHeadsKv为粒度切成 S1G 行将 ori_kv 与 cmp_kv 的 S2 轴以s2BaseSize_ 128为粒度切成 KV 块并统计每个 batch 的基础块数与尾块大小。模拟计费CalcCostInfo / CalcS1GCache根据 mask 模式CalcOriMaskMode/CalcCmpMaskMode将 0/3/4 分别映射为 No Mask、RightDownCausal、sliding window 的 pre/next token 范围计算每个 S1G 行在 S2 方向上的有效 token 区间CalcS2TokenRange结合稀疏 TopKGetOriTopkLength/GetCmpTopkLength在 mask_mode0 时取 topk_length 逐元素值否则取 topk 常量裁剪出实际负载再通过代价函数OriCalcCost/CmpCalcCost按“M 方向以 16 对齐、S2 方向以 64 对齐”的块粒度估算耗时COST_WEIGHT_M、COST_WEIGHT_S2为 M/S2 两个方向的权重系数。cmp_kv 部分还会用GetRevertS2Size将压缩后的长度乘以cmp_ratio并加上cmp_residual_kv余数还原为压缩前的 token 范围后再按比例折算负载。负载分配CalcSplitPlan依次尝试三种粒度的分配策略——整 batch 分配AssignByBatch、按 S1G 行分配AssignByRow、按 S2 块分配AssignByBlock仅 batch 一致性场景支持每轮以“剩余负载 / 剩余核数”得到costLimit并借助IsWithinTolerance容差因子FA_TOLERANCE_RATIO在接近目标时停止对无法按块均衡的场景用ForceAssign兜底。跨核行还会通过RecordFDInfo/SplitFD生成 AIV 核的 FD 归约任务划分按每个归约任务的fdS2SplitNum × fdMSize负载在冗余向量核上均分。生成 metadataGenMetadata将每个核的bN2End/gS1End/s2End起止索引与firstFdDataWorkspaceIdx写入faMetadata将 FD 任务信息写入fdMetadata未被使用的核 Core Enable 置 0。950 平台下若groupSize 64还会触发isSplitG_分支将 AIC 核数减半、每个核处理两个 G 组的任务以平衡负载。此外第一段 aclnn 接口aclnn_mixed_quant_sparse_flash_mla_metadata.cpp在执行前会读取平台信息GetCubeCoreNum/GetVectorCoreNum/GetSocLongVersion与确定性开关aclrtGetSysParamOpt(ACL_OPT_DETERMINISTIC)当确定性级别等于 3 时启用 batch 一致性模式isBatchConsistency此时 S2 方向改用规约级切分reductionBlockSize由BATCH_CONSISTENCY_MAX_REDUCTION_PARTS约束重新划分基本块保证跨 batch 的规约行为一致同时通过l0op::MixedQuantSparseFlashMlaMetadata声明见 mixed_quant_sparse_flash_mla_metadata.h下发到 AI CPU 执行。典型使用建议必须配套使用该算子输出的 metadata 是MixedQuantSparseFlashMla算子的输入之一切勿单独使用或与其他算子混用否则无法获得预期的负载均衡效果。layout 组合非 PA 场景下layout_q与layout_kv必须相同BSND 或 TNDPA_BBND 是仅 KV 侧支持的分页缓存布局常用于推理场景的 KV Cache。稀疏与非稀疏ori_topk/cmp_topk为 0 表示对应 KV 部分非稀疏不裁剪负载大于 0 表示稀疏稀疏 No Maskmask_mode0组合下必须逐 token 传入ori_topk_length/cmp_topk_length其余稀疏场景按 mask 规则与 topk 常量取最小有效 seqlen。压缩与余数cmp_ratio不为 1 且cmp_mask_mode3时cmp_residual_kv每个 batch 压缩后长度余数需满足 cmp_ratio必须传入否则无法正确还原 cmp_kv 的负载。seqlen 回退规则尽量显式传入seqused_*最精确BSND 下回退到max_seqlen_*此时必须传入真实 S1/S2/S3TND 下回退到cu_seqlens_*首元素必须为 0 且单调不减PA_BBND 下回退到*_topk_length。【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表