ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

CANN ops-transformer 算子 PyTorch 接入实战:stem_indexer 稀疏注意力前处理算子的 TorchNPU 扩展构建与调用

2026/9/19 12:40:11 拓冰建站 浏览量
CANN ops-transformer 算子 PyTorch 接入实战:stem_indexer 稀疏注意力前处理算子的 TorchNPU 扩展构建与调用 CANN ops-transformer 算子 PyTorch 接入实战stem_indexer 稀疏注意力前处理算子的 TorchNPU 扩展构建与调用【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer本篇技术指南围绕 CANN ops-transformer 仓库中experimental/attention/stem_indexer/torch_ops_extension目录展开系统讲解推理场景稀疏 Attention 前处理算子 StemIndexer 如何通过编译式NpuExtension桥接到 PyTorch注册为torch.ops.custom.npu_stem_indexer同时挂载到TorchNPU.npu_stem_indexer并底层调用 ACLNN 接口aclnnStemIndexer/aclnnStemIndexerMetadata。读完本文你将掌握该扩展的目录结构、构建安装流程、eager 单算子与 torch.compile 成图两种调用通路、接口与算子 IR 的对应关系以及如何在算子 pytest 框架中集成验证。一、StemIndexer 是什么稀疏 Attention 的块级打分与动态选块StemIndexer 是推理场景下稀疏 Attention 的前处理算子承担块级打分与动态选块职责。对于每个 Query Block算子基于qflat与kflat的相关性并叠加 Value 量值偏置vbiasOAMOutput-Aware Metric进行打分随后按 Position-Decay 动态 TopK 预算选出关键 Key Block输出sparse_indices与sparse_seq_len供下游 Block Sparse Attention 算子使用。完整的算子设计说明见 docs/StemIndexer.md。固定参数为stem_block_size128、stem_stride16因此代表向量聚合比例R B_s/T_s 8扁平化特征维度D_f T_s·D 2048原始 Head DimD 128。打分公式为S_{i,j} Q_i, K_j / R^2 V_j Q_i, K_j / 64 V_j其中1/R^2 1/64对每组内部 64 个 Token Pair 贡献进行归一化。最终输出集合按Sink Block → 普通 TopK 结果 → Window Block顺序拼接不执行额外去重。本文聚焦的是该算子的PyTorch 接入层——即torch_ops_extension目录它实现了算子从 Python 侧到 ACLNN C 接口的完整调用链。二、接入层目录结构与职责划分experimental/attention/stem_indexer/torch_ops_extension/ ├── setup.py # NpuExtension custom_ops.custom_ops_lib编译 csrc/*.cpp ├── build_and_install.sh # 构建 wheel 并 pip 安装 ├── README.md └── custom_ops/ ├── __init__.py # 导入 custom_ops_lib(.so) converter并挂载到 TorchNPU ├── csrc/ │ ├── ops_def_registration.cpp # TORCH_LIBRARY(custom): npu_stem_indexer(_metadata) schema │ ├── ops_common.h / ops_common.cpp # EXEC_NPU_CMD_V1 / ConvertType 基础设施拷贝 │ ├── npu_stem_indexer.cpp # NPU/Meta 前向实现 TORCH_LIBRARY_IMPL │ └── npu_stem_indexer_metadata.cpp └── converter/ ├── __init__.py ├── npu_stem_indexer.py # torch.compile (GE) 成图 converter └── npu_stem_indexer_metadata.py整体结构与参考实现quant_block_sparse_attn/torch_ops_extension对齐采用编译式NpuExtensioncustom_ops.custom_ops_libTORCH_LIBRARY/TORCH_LIBRARY_IMPLEXEC_NPU_CMD_V1的组合。其中csrc/ops_common.h、csrc/ops_common.cpp为该参考的整文件拷贝自包含提供 ACLNN 接口动态加载与参数类型转换的公共基础设施。三、前置条件接入层对运行环境有明确要求与仓库 README.md 一致Linux 操作系统Python 3.8GCC 9.4.0PyTorch 2.6.0 与匹配版本的TorchNPUAscend CANN Toolkit环境变量ASCEND_HOME_PATH已设置已部署aclnnStemIndexer/aclnnStemIndexerMetadata算子包运行时可经ASCEND_CUSTOM_OPP_PATH/ASCEND_OPP_PATH检索到libcust_opapi.so。算子本体的硬件适配情况见 docs/StemIndexer.md当前仅适配Ascend 950PR / Ascend 950DTA3、A2、Atlas 200I/500 A2 及其他训练/推理系列产品均不支持。在 op_host/stem_indexer_def.cpp 中算子定义仅执行AICore().AddConfig(ascend950, aicoreConfig)代码注释亦明确“当前仅适配 A5ascend950A2/A3 暂未适配不注册”。四、构建与安装wheel 与就地编译两种方式进入experimental/attention/stem_indexer/torch_ops_extension目录后提供两种构建方式# 方式一构建 wheel 并安装推荐 bash build_and_install.sh # 方式二就地编译.so 生成在 custom_ops/custom_ops_lib*.so python3 setup.py build_ext --inplacebuild_and_install.sh 的内部流程为清理历史build目录 →python3 setup.py build bdist_wheel构建 wheel → 在dist目录执行pip3 install *.whl -I强制安装。从 setup.py 可以看到编译细节使用torch_npu.utils.cpp_extension.NpuExtension扩展模块命名为custom_ops.custom_ops_lib源文件为custom_ops/csrc/*.cppglob 收集通过torch_npu安装根目录定位 ACL 头文件include/third_party/acl/inc并作为-I编译参数传入支持USE_NINJA1环境变量启用 ninja 加速默认关闭package_data显式包含*.py与*.so保证安装后 Python 侧from . import custom_ops_lib可加载编译产物。五、Python 侧注册与挂载机制custom_ops/__init__.py是整个接入层的入口其职责有三导入编译产物custom_ops_lib.so触发TORCH_LIBRARY(custom, m)中的 schema 注册导入converter子模块注册 torch.compile 成图所需的 GE converter遍历torch.ops.custom命名空间下所有算子逐个setattr(torch_npu, op_name, custom_op_func)挂载到torch_npu从而支持TorchNPU.npu_stem_indexer(...)调用形式若torch.ops.custom不存在则输出 warning 并降级为仅torch.ops.custom.xxx调用。schema 定义位于 csrc/ops_def_registration.cppTORCH_LIBRARY(custom, m) { m.def(npu_stem_indexer(Tensor qflat, Tensor kflat, Tensor vbias, Tensor q_seq_lens, Tensor kv_seq_lens, *, Tensor? num_prompt_tokensNone, Tensor? metadataNone, bool causalTrue, int stem_block_size128, int stem_stride16, float alpha1.0, int initial_blocks4, int window_size4, float k_block_num_rate_medium0.2, int k_block_num_bias_medium30, float k_block_num_rate_large0.1, int k_block_num_bias_large30, int topk_score_precision1) - (Tensor, Tensor)); m.def(npu_stem_indexer_metadata(Tensor q_seq_lens, Tensor kv_seq_lens, int q_heads, int kv_heads, *, bool causalTrue, int stem_block_size128, int dim_qkflat2048, int window_size4) - Tensor); }schema 设计遵循“必选张量 必选标量在前*之后为可选张量与带默认值的属性”的约定张量 dtype 与属性默认值对齐算子原型 op_host/stem_indexer_def.cpp 中Input/Attr的声明causal默认true、stem_block_size128、stem_stride16、alpha1.0、initial_blocks4、window_size4、k_block_num_rate_medium0.2、k_block_num_bias_medium30、k_block_num_rate_large0.1、k_block_num_bias_large30、topk_score_precision1。六、eager 调用单算子前向实现NPU/Meta 前向实现位于 csrc/npu_stem_indexer.cpp实现要点如下入参校验对qflat/kflat断言 4D、vbias断言 3Dqflat/kflat必须为 BFloat16vbias必须为 Float32kflat的 Kb 维与vbias一致qflat/kflat最后一维一致topk_score_precision仅允许 1UINT32或 2UINT16输出推导sparse_indices为 INT32[B, q_heads, Qb, Kb]sparse_seq_len为 INT32[B, q_heads, Qb]ACLNN 调用通过EXEC_NPU_CMD_V1(aclnnStemIndexer, ...)下发实参顺序严格按算子 IR 声明顺序输入 → 属性 → 输出与 Python schema 的形参顺序不同——schema 为便于调用将可选张量前置C 实现内部已按 IR 顺序重排设备注册TORCH_LIBRARY_IMPL(custom, PrivateUse1, m)注册 NPU 实现TORCH_LIBRARY_IMPL(custom, Meta, m)注册 Meta 实现仅做 shape/dtype 推导不真正计算。Metadata 接口 csrc/npu_stem_indexer_metadata.cpp 用于生成分核调度信息其输出容量按16 B·kv_heads·(36 72)·16计算并向上对齐到 4096 个 INT32 元素常量在源码中以STEM_INDEXER_METADATA_HEADER_SIZE16、STEM_INDEXER_METADATA_AIC_CORE_NUM36、STEM_INDEXER_METADATA_AIV_CORE_NUM72、STEM_INDEXER_METADATA_CORE_STRIDE16、STEM_INDEXER_METADATA_ALIGN_SIZE4096定义与文档中M AlignUp(16 C_max(3672)×16, 4096)、C_max B·N_k的容量公式一致。eager 单算子调用示例import torch, torch_npu import custom_ops # 注册 torch.ops.custom.npu_stem_indexer(_metadata) 并挂载到 torch_npu metadata torch.ops.custom.npu_stem_indexer_metadata( q_seq_lens, kv_seq_lens, q_heads, kv_heads, causalTrue, stem_block_size128, dim_qkflat2048, window_size4, ) sparse_indices, sparse_seq_len torch.ops.custom.npu_stem_indexer( qflat, kflat, vbias, q_seq_lens, kv_seq_lens, causalTrue, stem_block_size128, stem_stride16, alpha1.0, initial_blocks4, window_size4, k_block_num_rate_medium0.2, k_block_num_bias_medium30, k_block_num_rate_large0.1, k_block_num_bias_large30, topk_score_precision1, num_prompt_tokensnum_prompt_tokens, metadatametadata, ) # 也可TorchNPU.npu_stem_indexer(...)七、torch.compile 成图GE converterconverter目录下的 Python 文件为 torch.compileGE 图模式场景提供成图转换能力converter/npu_stem_indexer.py 通过register_fx_node_ge_converter(torch.ops.custom.npu_stem_indexer.default)注册将 FX 节点转换为torchair.ge.custom_op(StemIndexer, ...)输入为qflat/kflat/vbias/q_seq_lens/kv_seq_lens/num_prompt_tokens/metadata属性逐一映射causal→attr.Bool、整型属性→attr.Int、浮点属性→attr.Float输出为[sparse_indices, sparse_seq_len]converter/npu_stem_indexer_metadata.py 同理注册npu_stem_indexer_metadata生成 GE 节点StemIndexerMetadata。converter 中meta_outputs形参为固定写法若写错会影响 GE 节点的输出 dtype 与 shape 推导。由于接口“支持图模式”见 docs/StemIndexer.md 约束说明converter 是图模式下可用的前提。八、接口与 IR 对齐张量布局与可选张量语义接入层与算子原型的对应关系摘自 torch_ops_extension/README.md入参/属性/输出与算子原型 op_host/stem_indexer_def.cpp 一一对应EXEC_NPU_CMD_V1实参按 IR 声明顺序传入输入→属性→输出Python schema 为便于调用将必选张量前置C 实现内部已按 IR 顺序重排qflat/kflat为 BF16[B, N, Qb/Kb, stem_stride*D]vbias为 FP32[B, kv_heads, Kb]其余q_seq_lens/kv_seq_lens/num_prompt_tokens/metadata为 INT32metadata长度按16 B * kv_heads * (36 72) * 16计算并向上对齐到 4096 个 INT32 元素输出sparse_indices/sparse_seq_len为 INT32num_prompt_tokens/metadata为可选张量c10::optionalat::Tensor两条调用通路eager 与 graph均将缺省状态传递给 OpHost。未提供num_prompt_tokens时TilingData 记录复用标志Kernel 使用kv_seq_lens未提供metadata时OpHost Tiling 按当前计算要求返回明确的参数错误。关于可选张量的缺省行为测试框架 tests/pytest/README.md 也有印证正例中保持num_prompt_tokens kv_seq_lens缺省时由 OpHost 通过 TilingData 通知 Kernel 复用kv_seq_lensmetadata虽然在接口层声明为可选但当前主算子计算必须传入有效 Metadata缺省时会在 Tiling 阶段返回参数错误。九、完整参数说明与约束以下参数表完整继承自 docs/StemIndexer.md维度符号$B$ 为 Batch$N_q$/$N_k$ 为 Query/KV Head 数$Q_{\max}$/$K_{\max}$ 为最大 Query/KV Block 数$D_f T_s·D 2048$。参数名输入/输出描述使用说明数据类型qflat输入Q 侧块级压缩表示Kernel 按连续布局处理标准执行路径声明了 AutoContiguousBF16kflat输入K 侧块级压缩表示最后一维必须与qflat一致BF16vbias输入Value 量值偏置OAM 项Batch、KV Head 和 KV Block 维必须与kflat一致FLOAT32q_seq_lens输入每个 Batch 的 Query 有效 Token 数非 Block 数按 $Q_b\lceil\text{q_seq_lens}[b]/B_s\rceil$ 计算有效 Query Block 数INT32kv_seq_lens输入每个 Batch 的 KV 有效 Token 数非 Block 数按 $K_b\lceil\text{kv_seq_lens}[b]/B_s\rceil$ 计算有效 KV Block 数INT32num_prompt_tokens可选输入每个 Batch 的 Prompt Token 数用于 Position-Decay 动态 TopK 预算分档未传入时复用kv_seq_lensINT32metadata可选输入分核调度信息接口声明为可选但当前计算必须使用有效 Metadata未传入时 Tiling 返回参数错误INT32causal输入属性是否采用 Right-down Causal 语义默认trueBOOLstem_block_size输入属性一个 Stem Block 包含的原始 Token 数当前仅支持 128默认 128INT64stem_stride输入属性Stem Block 内部的分组数/聚合 Stride当前仅支持 16默认 16INT64alpha输入属性动态 TopK 预算随 Query 位置的衰减程度$K_e K_s·\alpha$取值范围 $(0,1]$默认 1.0越小衰减越强FLOAT32initial_blocks输入属性开头固定保留的 Sink Block 数当前仅支持 4默认 4INT64window_size输入属性末尾固定保留的 Window Block 数当前仅支持 4默认 4INT64k_block_num_rate_medium输入属性中等长度 Prompt 的 TopK 预算系数当前仅支持 0.2默认 0.2FLOAT32k_block_num_bias_medium输入属性中等长度 Prompt 的 TopK 预算偏置当前仅支持 30默认 30INT64k_block_num_rate_large输入属性长 Prompt 的 TopK 预算系数当前仅支持 0.1默认 0.1FLOAT32k_block_num_bias_large输入属性长 Prompt 的 TopK 预算偏置当前仅支持 30默认 30INT64topk_score_precision输入属性TopK 内部可排序 Score 的存储精度1 表示 UINT322 表示 UINT16默认 1不改变输出 Tensor 数据类型INT64sparse_indices输出选中的 Key Block 逻辑索引每行仅前sparse_seq_len项有效尾部无效区填充 -1INT32sparse_seq_len输出每个 Query Block 对应的有效 Key Block 数量无有效 Query/KV 任务时对应值为 0INT32关键约束速览详见 docs/StemIndexer.md支持图模式仅适配 Ascend 950PR/950DT$B$ 取值范围 $[1,65536]$$N_q$ 仅支持 32 或 64$N_k$ 仅支持 2、4 或 8且 $N_q \bmod N_k 0$metadata必须由与主算子相同的q_seq_lens、kv_seq_lens、$N_q$、$N_k$、causal、stem_block_size、$D_f$ 和window_size生成Sink Block 和 Window Block 不参与普通 TopK 候选最终按 Sink → TopK → Window 顺序拼接不执行去重sparse_indices有效前缀元素为 Key Block 逻辑索引第 $b$ 个 Batch 取值 $[0, K_b-1]$调用者不应依赖有效索引的排列顺序。十、动态 TopK 预算算法动态 TopK 预算Position-Decay的计算逻辑是选块行为的核心公式如下$B_s128$Prompt Token 数转 Block 数$P_b \lceil L_b^p / B_s \rceil$初始预算分档K_s P_b , P_b 56 K_s floor(0.2·P_b 30), 56 ≤ P_b 160 K_s floor(0.1·P_b 30), P_b ≥ 160衰减终点 $K_e \alpha·K_s$对第 $i$ 个 Query Block按 Right-down 对齐位置 $p_i i K_b - Q_b$ 线性插值并 clamp 到 $[1, K_s]$实际普通 TopK 数量再限制为不超过 256。alpha1.0表示不衰减值越小序列后部选择的 Key Block 越少。该分段边界55/56/159/160 blocks在 tests/pytest/README.md 中被列为显式测试覆盖点。十一、Ascend 950PR/950DT 完整调用示例以下示例完整继承自 docs/StemIndexer.md展示了从构造输入、生成 metadata 到执行选块的完整流程import torch import torch_npu import custom_ops batch_size 4 q_heads 64 kv_heads 8 d 128 stem_block_size 128 stem_stride 16 q_block_num 16 k_block_num 64 q_seq_len q_block_num * stem_block_size kv_seq_len k_block_num * stem_block_size prompt_token_num kv_seq_len torch.manual_seed(0) qflat torch.randn( batch_size, q_heads, q_block_num, stem_stride * d, dtypetorch.bfloat16, ).npu() kflat torch.randn( batch_size, kv_heads, k_block_num, stem_stride * d, dtypetorch.bfloat16, ).npu() vbias torch.randn( batch_size, kv_heads, k_block_num, dtypetorch.float32, ).npu() q_seq_lens torch.full((batch_size,), q_seq_len, dtypetorch.int32).npu() kv_seq_lens torch.full((batch_size,), kv_seq_len, dtypetorch.int32).npu() num_prompt_tokens torch.full((batch_size,), prompt_token_num, dtypetorch.int32).npu() # 1. 生成与主算子输入匹配的分核调度 Metadata。 metadata torch.ops.custom.npu_stem_indexer_metadata( q_seq_lens, kv_seq_lens, q_heads, kv_heads, causalTrue, stem_block_sizestem_block_size, dim_qkflatstem_stride * d, window_size4, ) # 2. 执行块级打分与动态选块。 sparse_indices, sparse_seq_len torch.ops.custom.npu_stem_indexer( qflat, kflat, vbias, q_seq_lens, kv_seq_lens, num_prompt_tokensnum_prompt_tokens, metadatametadata, causalTrue, stem_block_sizestem_block_size, stem_stridestem_stride, alpha1.0, initial_blocks4, window_size4, k_block_num_rate_medium0.2, k_block_num_bias_medium30, k_block_num_rate_large0.1, k_block_num_bias_large30, topk_score_precision1, )十二、与 pytest 测试框架集成扩展可通过环境变量STEM_INDEXER_CUSTOM_OPS_PATH或默认相对路径stem_indexer/torch_ops_extension被测试检索既支持 globcustom_ops_lib*.sotorch.ops.load_library也支持 execcustom_ops/__init__.py。pytest 测试框架 的验证思路为CPU 侧按设计实现 goldenNPU 侧先经npu_stem_indexer_metadata生成分核信息再经npu_stem_indexer获取实际结果支持 eager 与 graphtorch.compiletorchair编译为 aclgraph两种模式STEM_INDEXER_MODE切换默认 eager结果比对仅比较sparse_seq_len与sparse_indices有效前缀尾部未定义区域不校验。执行方式在tests/pytest目录下bash test_run.sh single # single (eager) bash test_run.sh single_graph # single (graph) bash test_run.sh batch # batch (eager) bash test_run.sh batch_graph # batch (graph)指定用例运行STEM_INDEXER_CASE_ID支持逗号分隔多个 case_idSTEM_INDEXER_CASE_IDSI_WB_001,SI_WB_002 python3 -m pytest test_stem_indexer_single.py STEM_INDEXER_CASE_IDSI_WB_001,SI_WB_002 python3 -m pytest test_stem_indexer_batch.py测试覆盖点包括q/kv 尾块与空序列边界、Sink/Window 的重叠与裁剪、causal 与 non-causal 路径、TPDalpha无衰减/普通衰减/强衰减、动态 TopK small/medium/large 分段边界55/56/159/160 blocks、GQA 组合q_heads32/64 ×kv_heads2/4/8 共 6 种合法组合、多 batch 变长与 prefill/decode 混合、32K/64K/128K/256K/1M token 长序列量级、OAMvbias影响选块与scoreScale1/64路径以及topk_score_precision的 uint32SI_WB_001~SI_WB_100与 uint16SI_WB_101~SI_WB_150镜像两条路径。十三、底层基础设施EXEC_NPU_CMD_V1 与动态库加载csrc/ops_common.h自包含拷贝提供了接入层运行时的关键支撑理解它能更好地把握整条调用链ACLNN 接口动态加载GetOpApiFuncAddr按ASCEND_CUSTOM_OPP_PATH→ASCEND_OPP_PATH/vendors读取config.ini的load_priority→ 各 feature 库libopapi_math.so/libopapi_nn.so/libopapi_cv.so/libopapi_transformer.so/libopapi_legacy.so→libopapi.so的顺序查找aclnnStemIndexer等接口符号自定义算子优先从libcust_opapi.so解析类型转换ConvertType将at::Tensor/at::Scalar/at::IntArrayRef/c10::optionalat::Tensor等 PyTorch 类型转换为 ACL 的aclTensor/aclScalar/aclIntArraydtype 经ConvertScalarTypeToAclDataType映射如BFloat16→ACL_BF16、Float→ACL_FLOAT、Int→ACL_INT32执行宏EXEC_NPU_CMD_V1(aclnn_api, ...)依次完成xxxGetWorkspaceSize查询 workspace、按需分配 workspace Tensor、在当前 NPU 流c10_npu::getCurrentNPUStream()上异步执行 ACLNN 接口并在执行后释放转换产物。正是这套基础设施使得npu_stem_indexer/npu_stem_indexer_metadata两个自定义算子能够以“与原生算子一致”的方式在 NPU 流上排队执行。十四、小结torch_ops_extension是 StemIndexer 算子从 Python 侧触达 NPU 的唯一入口其“schema 定义 → NPU/Meta 双实现 → converter 成图”三层结构清晰可复用ops_def_registration.cpp声明接口契约npu_stem_indexer.cpp/npu_stem_indexer_metadata.cpp实现 eager 通路converter/支撑 torch.compile 图模式ops_common屏蔽了 ACLNN 动态加载与类型转换的复杂度。如需在自有业务中接入 StemIndexer推荐路径是按本文第四节完成构建安装 → 按第十一节组织输入并先调用 metadata 接口 → 再调用主算子若需验证正确性可直接复用 pytest 框架 的 golden 比对流程。相关参考算子完整设计docs/StemIndexer.md算子原型定义op_host/stem_indexer_def.cpp接入层 READMEtorch_ops_extension/README.mdpytest 验证框架tests/pytest/README.md【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考