
算子库人工智能深度学习Ascend【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址https://gitcode.com/cann/ops-transformer点击查看免费下载本篇技术指南以 CANN ops-transformer 开源仓库中的 torchapi_msa_index_score.md 为主线系统讲解 MSAMiniMax Sparse AttentionIndex Branch 中 block score 算子的 PyTorch 封装cann_ops_transformer.msa_index_score从计算公式、函数原型、完整参数语义到 PageAttentionBBND/BNBD与 packed TND 两种实际调用路径并深入 tiling 实现 与 CPU golden 参考实现 佐证关键行为。读完本文你将能够正确地在 Atlas A2/A3 与 Ascend 950 上构造合法入参、选择量化/FP8 模式、理解local_mask的强制选块机制并掌握其底层 maxpool 归约与核调度原理。一、产品支持情况msa_index_score在以下 NPU 平台上受支持见 算子 README产品是否支持Atlas A2 训练系列产品 / Atlas A2 推理系列产品支持Atlas A3 训练系列产品 / Atlas A3 推理系列产品支持Ascend 950PR / Ascend 950DT支持其中 Ascend 950 相比 A2/A3 额外支持三种 FP8 输入类型详见“量化与 FP8 支持”一节并在核实现上走独立的arch35路径见 op_kernel/arch35A2/A3 则使用arch22路径。二、功能说明与计算公式该算子封装底层aclnnMsaIndexScoreaclnn 接口详见 aclnnMsaIndexScore.md核心职责是对每个 query token 与每个 KV sparse block计算该 block 的重要性分数score作为 MSA Index Branch 中后续 TopK 选块阶段的输入。Prefill 与 Decode 由同一接口承载。其数学本质是对每个 sparse block长度为block_size内的 KV token 做一次QKᵀ矩阵乘后按 KV 维取最大值Maxpool得到逐 block 的标量分数。完整公式为$$ score Maxpool[(scale\cdot)Q_{idx}K_{idx}^{T}atten_mask]local_mask $$其中各分量含义如下Q_idx/K_idxIndex Branch 的 query / key即用于选块的轻量索引张量非注意力主路径的完整 Q/Kscaleint8 量化场景下的反量化系数非量化不传atten_mask因果可见性掩码仅在sparse_mode3时使用local_mask由start_loc、init_blocks、local_blocks共同生成用于在 Maxpool 之后对若干“必须保留”的 block 写入强制高分保证下游 TopK 一定选中它们。local_mask的构造规则与 aclnnMsaIndexScore.md 及 msa_index_score_common.h 中的常量一致逻辑 block[0, init_blocks)写入高分1e30对应常量MSA_LOCAL_SCORE_INIT窗口[max(0, start_loc 1 - local_blocks), start_loc]写入高分1e29对应常量MSA_LOCAL_SCORE_LOCAL覆盖同位置的 init 高分默认值init_blocks0、local_blocks1对齐 MiniMax HF 常见配置当两者均为 0 时跳过local_mask。与 Triton raw score kernel 对齐对比时需要置 0。需要注意本算子输出止于 block score不包含 TopKTopK 是 Index Branch 下游的另一阶段。三、函数原型cann_ops_transformer.msa_index_score( query, key, start_loc, *, block_tableNone, scaleNone, atten_maskNone, actual_seq_qlenNone, actual_seq_klenNone, layout_keyBBND, sparse_mode3, init_blocks0, local_blocks1, ) - Tensor该接口由 torch_extension/msa_index_score.py 定义。其底层注册的 torch schema 与 Python 默认值完全对应layout_key默认BBND、sparse_mode默认3、init_blocks默认0、local_blocks默认1该文件中的DEFAULT_SPARSE_MODE/DEFAULT_INIT_BLOCKS/DEFAULT_LOCAL_BLOCKS/DEFAULT_LAYOUT_KEY常量。从源码可见两点实现细节layout_key支持TND、BBND、BNBD三种取值大小写不敏感内部统一upper()传入非法值会抛出ValueError当layout_keyBBND且 key 是 3 维[NP, P, D]时即N21的紧凑形式Python 层会自动unsqueeze(2)补成 4 维[NP, P, 1, D]方便用户少写一维。四、参数说明下表完整整理自关联文档torchapi_msa_index_score.md并结合 算子原型注册 补充了 dtype 组合细节。先约定符号B 为 batchS1/S2 为 query/key 侧序列长度T1/T2 为所有 batch 序列长度的累加和packed token 数N1/N2 为 query/key 侧头数D 为单头维度。PageAttention 场景下block_num为物理 block 总数、block_size为每个 block 的 token 数、maxBlockNumPerSeq为每个 batch 最大逻辑 block 数、M_b ⌈S2/block_size⌉为逻辑 block 总数。参数名可选/必选描述dtypeshapequery必选Q_idxTND 布局fp16/bf16950 另支持 fp8_e4m3fn / fp8_e5m2 / hifloat8[T1, N1, D]key必选K_idxPA BBND/BNBD 或 TND。A2/A3 与 950 上 PA 允许 dim0物理 page非连续同 query 或 int8int8 配 fp16 query[NP, P, N2, D]/[NP, N2, P, D]/[T2, N2, D]start_loc必选当前 query 所在逻辑 block 索引非 token 前缀用于生成local_maskint32[B]block_tablePA 必选逻辑 block → 物理 page 映射表TND 不传int32[B, MB]scale量化必选int8 反量化系数非量化 / FP8 必须为空float32PA[NP, N2, P]或 BBN 序TND[T2, N2]N21 时可[T2]atten_maskmode3 必选压缩下三角因果模板1 表示该位不参与计算0 表示参与int8[2048, 2048]actual_seq_qlenTND query 必选query 各 batch 长度前缀和单调不减int32[B1]actual_seq_klen必选PA 为各请求可见S2TND 为 key 前缀和int32[B]/[B1]layout_key可选key 布局TND/BBND/BNBD默认BBNDstr-sparse_mode可选0 或 3默认 3int-init_blocks可选local_mask头部强制块数默认 0int-local_blocks可选local_mask局部窗口长度默认 1int-补充说明几点参数语义细节sparse_mode0表示 defaultMask无因果atten_mask不传3表示 rightDownCausal右顶点划分的下三角因果必须传[2048, 2048]的atten_mask。device 侧按 rightDownCausal 解析可见窗口与 LightningIndexer 一致不逐元素加载 mask 模板atten_mask主要用于 host 侧合法性校验。该语义在 msa_index_score_common.h 中由MSA_SPARSE_MODE_DEFAULT 0、MSA_SPARSE_MODE_RIGHT_DOWN 3两个常量固化。block_table第二维可以大于实际 KV 逻辑 block 数例如 vLLM 预分配的宽表score 末维取RoundUp(宽表宽度, 16)。q_len/kv_len允许为 0含整 batch对应请求跳过 QK 计算空 KV 的 score 填-inf整 batchT10时 host 侧SetBlockDim(1)见 tiling 实现。五、输出说明输出为逐 block 的重要性分数张量[N1, T1, RoundUp(MB, 16)] dtype float32其中MB为最大逻辑 block 数PA 场景取block_table第二维TND 场景取各请求 block 数最大值末维按16 对齐常量MSA_SCORE_STRIDE_ALIGN 16。输出恒为 float32Cube 以 fp32 累加下游 TopK 直接消费 fp32见 infershape 实现 中InferDataTypeMsaIndexScore。对齐宽度的来源可从 inferShape 与 Python meta 实现确认TND 场景若所有请求kv_len0maxBlocks会被钳到 1再 16 对齐得到末维 16此时核内跳过 QK、把 score 填-inf。六、调用示例6.1 PageAttention BBND默认布局import torch import torch_npu import cann_ops_transformer T1, N1, N2, D, P 32, 8, 1, 128, 128 B, NP, MB 1, 8, 2 query torch.randn(T1, N1, D, dtypetorch.float16).npu() key_bbnd torch.randn(NP, P, N2, D, dtypetorch.float16).npu() # BNBD 与 BBND 仅 N/P 轴对调key_bnbd key_bbnd.permute(0, 2, 1, 3).contiguous() block_table torch.arange(B * MB, dtypetorch.int32).view(B, MB).npu() actual_seq_qlen torch.tensor([0, T1], dtypetorch.int32).npu() actual_seq_klen torch.tensor([256], dtypetorch.int32).npu() start_loc torch.tensor([1], dtypetorch.int32).npu() atten_mask torch.zeros(2048, 2048, dtypetorch.int8).npu() score cann_ops_transformer.msa_index_score( query, key_bbnd, start_loc, block_tableblock_table, atten_maskatten_mask, actual_seq_qlenactual_seq_qlen, actual_seq_klenactual_seq_klen, layout_keyBBND)要点query是 packed 的 TND 布局[T1, N1, D]T1 为所有 batch query token 之和key_bbnd是 PageAttention 的物理 page 布局[NP, P, N2, D]其中NP8个物理 page、每 page 恰好一个 sparse blockP block_size 128start_loc[1]表示该请求当前 query 位于逻辑 block 1与默认local_blocks1组合意味着逻辑 block 1 会被local_mask强制高分从 CPU golden 参考实现 可以看到golden 对每个 block 取block_table[b, blk]指向的物理 page再对 block 内可见 token 求q k_pageᵀ的逐行最大值即与上述公式严格对应。6.2 PageAttention BNBD仅需把 key 调整为[NP, N2, P, D]并指定layout_keyBNBD其余入参完全一致key_bnbd key_bbnd.permute(0, 2, 1, 3).contiguous() # [NP, N2, P, D] score cann_ops_transformer.msa_index_score( query, key_bnbd, start_loc, block_tableblock_table, atten_maskatten_mask, actual_seq_qlenactual_seq_qlen, actual_seq_klenactual_seq_klen, layout_keyBNBD)6.3 TND packed keyTND 场景不传block_tableactual_seq_klen传[B1]前缀和T2 256 key_tnd torch.randn(T2, N2, D, dtypetorch.float16).npu() actual_seq_klen_tnd torch.tensor([0, T2], dtypetorch.int32).npu() score_tnd cann_ops_transformer.msa_index_score( query, key_tnd, start_loc, atten_maskatten_mask, actual_seq_qlenactual_seq_qlen, actual_seq_klenactual_seq_klen_tnd, layout_keyTND)TND 布局下 key 是连续 packed 的[T2, N2, D]kernel 内部按actual_seq_klen前缀和切出各请求的 KV 区间并切块block_size128计算。若希望从 aclnn C 侧构造完全等价的三布局调用BBND/BNBD/TND 的差异代码可参考 examples/test_aclnn_msa_index_score.cpp 中的L1-bnbd*、L1-tnd*、L0-tnd-tiny用例。七、量化与 FP8 支持7.1 int8 量化路径适用条件key为 int8query必须为 fp16scale必选float32。PA 场景 scale shape 为[NP, N2, P]或 BBN 序TND 场景为[T2, N2]N21 时允许[T2]。数值路径score Maxpool[scale · QKᵀ]即 kernel 前融合反量化——由 AIV 先将 int8 K 乘 scale 转成 fp32/fp16再交给 Cube 做 mmad。对应 tiling 中isQuant分支int8 路径需要额外的 K cast scratch 空间A2/A3 上为四槽轮转MSA_K_SCRATCH_STAGES_A2 4950 上为单页握手见 msa_index_score_common.h。测试矩阵中的L0-int8-dequant-trace、L1-int8-dequant、L1-bnbd-int8、L1-stride-int8等用例覆盖该路径见 tests/README.md。7.2 FP8 路径仅 Ascend 950950 额外支持 query/key同型的torch.float8_e4m3fn/torch.float8_e5m2/torch_npu.hifloat8三种 FP8 均不得传scale无 int8 反量化语义。核侧为 Cube原生 FP8计算TilingKey 4/5/6 对应 hifloat8 / fp8_e5m2 / fp8_e4m3fn见 msa_index_score.cpp 所引用的 TilingKey 宏与 msa_index_score_common.h 的MSA_TILING_KEY_*定义不做 Cast→fp16 的降精度中转。已知限制torch_npu.hifloat8.npu()当前会产生非法 device id测试脚本会跳过该用例但 kernel 与 aclnn 接口均已注册 HIFLOAT8见 torch_extension/msa_index_score.py 与 msa_index_score_def.cpp 的 950 配置。八、关键实现原理源码级8.1 key 布局与连续性校验layout_key在 host 侧显式指定不从 shape 推断tiling 的ParseLayoutKeyAttr。PA key 允许 dim0物理 page非连续tiling 通过GetRequiredInputStride/GetInputStride读取 dim0 元素 stride写入strideKvBlock供 kernel 按 stride 寻址对应key | gap | key | ...的存储形态。这一能力对 vLLM 类 paged 缓存中因 preempt/append 导致的物理页不紧凑场景至关重要用户只需保持 key 为 view、不要调用.contiguous()。除 dim0 外size1 的非首轴 stride 必须与紧凑布局一致否则 tiling 拒绝ValidateKeyInnerAxesContiguous对齐 MlaProlog 的GetCacheStride0逻辑size1 的轴不参与寻址因此跳过比较例如 BBND↔BNBD 且 N21 时 permute 后 dim1 stride 可能不等于紧凑值属合法。TND key 不允许非连续dim0 stride 必须等于紧凑值否则报错。scale仍按逻辑 page 紧凑布局处理。8.2 Maxpool 归约与 score 对齐每个 sparse block 内部block_size128个 token 的 QK 分数先在 CubeMmad上完成[M, 128]的矩阵乘再由 AIV 沿 KV 维做 max 归约归约结果每行 8 个 fp32MSA_BLOCKS_PER_STILE8凑成 32B 对齐写回最后得到逐 block 的标量分数并写入按 16 对齐的 score 末维。空 KV 的不可见 block 填MSA_FILL_VALUE -3.4028e38F即-inf使其在下游 TopK 中必然落选local_mask的高分常量1e30/1e29也在 msa_index_score_common.h 中定义。8.3 核调度短 decode 不打满 AICA2/A3 与 950 的 host tiling 按估计 M-task 数启动 MIX1 AIC : 2 AIV避免短 decode 时打满空核单 batchB1时估计任务数为CeilDiv(T1·N1, 128)多 batch 时用 packed 行数上界加 B 的修正EstPackedMTasks中packedM B上界并截到 AIC 总数整 batchq_len0时直接SetBlockDim(1)。Ascend 950 在 M-task 填不满 AIC 且可见 KV stile 1 时再按 KV stile 沿 S 维度切分kvChunks把短 M 长 S 的 decode 场景也打满 MIX宽表短 KV 仍只起 1 个 MIX对应 tiling 实现 的EstKvChunks/EstLaunchAic。8.4 宽 block_table 的 256 列滑窗950 的 C2UBC_to_UB路径中score 暂存单窗 256 列当 score 末维RoundUp(width, 16) 256时按 256 列滑窗 flush并在后续位置补写-inf不能关闭 stage。这解释了为何宽block_table如 257、275 列在 950 上依然正确L0-wide-table-257、L0-decode-q4-kv4-table275、L0-decode-q4-kv275等用例专门覆盖见 tests/README.md。8.5 InferShape 的输出推导infershape 实现 中PA 场景maxBlocks直接取block_table第二维TND 场景优先用actual_seq_klen前缀和差分计算各请求 block 数取最大值host 不可读时退化为⌈T2/128⌉。输出末维恒为RoundUp(maxBlocks, 16)。九、编译、运行与精度验证9.1 编译与运行示例Ascend 950cd /path/to/ops-transformer bash build.sh --pkg --socascend950 --opsmsa_index_score -j32 bash ./build_out/cann-ops-transformer-custom_linux-x86_64.run --quiet --install-path/path/to/msa_opp/ source /path/to/msa_opp/vendors/custom_transformer/bin/set_env.bash export ASCEND_CUSTOM_OPP_PATH/path/to/msa_opp/vendors/custom_transformer # aclnn example 必须带 --socascend950否则默认 910b bash build.sh --run_example msa_index_score eager cust --vendor_namecustom --socascend950 # 通过末行 [PASS]: 50/50 cases passedA2/A3 平台把--soc换成ascend910b期望结果[PASS]: 40/40 cases passed自动跳过 10 条 FP8。完整步骤与用例矩阵详见 tests/README.md。9.2 用例矩阵与判定标准默认测试矩阵覆盖fp16/bf16/int8 三 dtype ×BBND/BNBD/TND三布局、mixed-batch pad部分请求q_len0/kv_len0含整 batch 全空、PA key dim0 stridegap2 下毒校验、宽block_table257/275、短 decodeq4-kv4、q4-kv275、投机解码q_len1、kv4096 长序列多 S-tile、极小 kv 尾填充等950 额外含 10 条 FP8D128。容差fp16/bf16/int8 为1e-3FP8 为2e-2。端到端精度通过内置 CPU golden 自验证test_aclnn_msa_index_score.cppPython 参考实现见 tests/golden/msa_index_score_golden.py。判定标准见 tests/README.md填充位不可见 block两侧同为-inflocal_mask强制高分两侧同为≥ 1e28有效位atol/rtol 1e-3error_ratio ≤ 1e-3950 FP8 为2e-2。十、常见误区与注意事项不要对 PA key 调用.contiguous()dim0 非连续是受支持特性contiguous()反而会破坏与 paged cache 的零拷贝语义但 TND key 必须连续。layout_key必须显式指定且与 key shape 一致不能仅凭维度推断传错会直接报非法输入。sparse_mode3必须传[2048, 2048]的atten_masksparse_mode0则必须不传——host 侧会严格校验。int8 量化必须同时满足queryfp16、keyint8、scale必传且为 float32非量化/FP8 传scale会报错。init_blocks/local_blocks不能超过逻辑 block 数PA 为block_table第二维TND 为 score 末维对齐宽度对比 Triton raw score 时两者都置 0。start_loc是逻辑 block 索引而非 token 偏移与actual_seq_qlen前缀和是两种完全不同的量纲。950 上 FP8 的 query 与 key 必须同型且三种 FP8 均不支持scale。相关配套资料算子总览 README、aclnn 两段式接口文档、C 端到端示例、测试说明。赞分享算子库人工智能深度学习Ascend【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址https://gitcode.com/cann/ops-transformer点击查看免费下载相关推荐CANN ops-transformer aclnnMsaIndexScore 算子深度解析MSA Index Branch 逐 block 重要性分数的两段式接口与 NPU 实现CANN ops transformer aclnnMsaIndexScore 算子深度解析MSA Index Branch 逐 block 重要性分数的两段算子库人工智能深度学习AscendCANN ops-transformer 算子解析DistributeBarrierExtend 全卡同步算子原理与 aclnn 调用实践CANN ops transformer 算子解析DistributeBarrierExtend 全卡同步算子原理与 aclnn 调用实践 本技术指南围绕 C算子库人工智能深度学习AscendCANN ops-transformer 算子指南aclnnQuantFlashAttentionScoreGrad 融合量化注意力 Score 反向计算原理与两段式调用实践CANN ops transformer 算子指南aclnnQuantFlashAttentionScoreGrad 融合量化注意力 Score 反向计算原理算子库人工智能深度学习Ascend上一篇GitHub Pages重定向技术实战5分钟构建FOSSASIA照片分享网站下一篇forever源码中的单元测试并行执行提高测试速度创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考