PyPTO Tensor.topk 算子详解:在 CANN 昇腾平台上获取前 k 个最值及其索引
发布时间:2026/9/19 4:52:51来源:尧图网络
PyPTO Tensor.topk 算子详解在 CANN 昇腾平台上获取前 k 个最值及其索引【免费下载链接】pyptoPyPTO发音: pai p-t-oParallel Tensor/Tile Operation编程范式。项目地址: https://gitcode.com/cann/pypto本文围绕 CANN PyPTOParallel Tensor/Tile Operation 编程范式中Tensor.topk的完整使用链路展开从 Python 侧Tensor方法、pypto.topk顶层接口到算子参数语义、两种 TopK 算法MERGE_SORT与RADIX_SELECT的选型与 TileShape 切分配置并结合仓库源码说明其底层实现原理。读者读完后可以准确配置k、dim、largest、algo参数理解 TileShape 与 UB 内存约束并在 Atlas A2/A3 训练与推理系列产品及 Ascend 950 系列上正确编写 topk 程序。产品支持情况Tensor.topk方法在当前仓库所支持的昇腾产品上均可使用Ascend 950PR/Ascend 950DT支持Atlas A3 训练系列产品/Atlas A3 推理系列产品支持Atlas A2 训练系列产品/Atlas A2 推理系列产品支持需要特别说明的是上述“支持”指Tensor.topk方法本身及默认算法MERGE_SORT的支持范围RADIX_SELECT算法在不同型号上的支持度存在差异详见本文算法选型与约束说明一节。功能概述沿尾轴取前 k 个最值Tensor.topk用于获取张量最后一个维度的前 k 个最大值或最小值及其对应的索引如果输入是向量则在向量中找到前 k 个最大值或最小值及其对应的索引如果输入是矩阵则沿最后一个维度计算每一行中前 k 个最大值或最小值及其对应的索引。以 Shape 为(4, 32)的二维矩阵为例k 设置为 1 时输出为[[32] [32] [32] [32]]即取出了每一行的最大值从仓库实现来看该接口的方法原型位于 python/pypto/tensor.py其内部直接将k、dim、largest透传给顶层算子接口pypto.topk顶层算子的完整语义定义位于 docs/zh/api/tensor_api/operation/pypto-topk.md。函数原型Tensor.topk作为张量成员方法的原型如下topk(self, k: int, dim: Optional[int] None, largest: bool True) - Tuple[Tensor, Tensor]其底层对应的顶层算子接口 pypto.topk 的原型为topk(input: Tensor, k: int, dim: Optional[int] None, largest: bool True, algo: TopKAlgo TopKAlgo.MERGE_SORT) - Tuple[Tensor, Tensor]对比可见两种调用方式的关系为Tensor.topk(k, dim, largest)是便捷方法不暴露algo参数始终使用默认算法需要显式控制算法如选择RADIX_SELECT以获得 O(n) 复杂度时应直接调用pypto.topk(input, k, dim, largest, algo)。在 Python 侧两者最终都汇聚到 python/pypto/op/comparison.py 中由op_wrapper装饰的topk函数该函数将参数归一化dim is None时按-1处理后调用底层pypto_impl.TopKreturn pypto_impl.TopK(input, k, (-1 if dim is None else dim), largest, algo)pypto_impl.TopK进一步绑定到 C 侧npu::tile_fwk::TopK绑定关系见 python/src/bindings/operation.cpp其参数缺省值islargest true、algo TopKAlgo::MERGE_SORT与文档原型保持一致。参数说明参数名输入/输出说明input输入源操作数。支持的类型为Tensor。Tensor支持的数据类型为- MERGE_SORTDT_FP32。- RADIX_SELECTDT_BF16DT_FP16DT_FP32DT_INT64DT_UINT64DT_INT32DT_UINT32DT_INT16DT_UINT16DT_INT8DT_UINT8。不支持空TensorShape仅支持1-4维Shape Size不大于2147483647即INT32_MAX。k输入返回元素的数量。k的大小应该满足1 k input.shape[dim]。dim输入指定排序的维度。目前仅支持按最后一个维度排序即dim -1或dim input.shape.size() - 1。largest输入如果为True返回最大元素。如果为False返回最小元素。algo输入算法枚举类型用以控制TopK计算的流程具体定义为TopKAlgo。默认为MERGE_SORT归并排序算法。几点实现层面的补充说明数据类型校验在 C 侧TopK入口完成见 framework/src/interface/operation/vector/topk.cppMERGE_SORT仅接受DT_FP32RADIX_SELECT接受列表中的 11 种数据类型维度校验限制输入 Shape 为 14 维并强制要求axis len - 1 || axis -1否则报错TopK only supports the last axisRADIX_SELECT算法在非DAV_3510架构上会直接报错When TopK using radix select algo, only DAV_3510 architecture is supported.这与文档约束中“RADIX_SELECT 仅在 Ascend 950 系列支持”的说明一致输出索引张量固定为DT_INT32输出值张量与输入数据类型一致输出 Shape 在dim轴上被替换为k。返回值说明返回一个命名元组(values, indices)其中包含input在指定维度dim下每行中最大或最小的 k 个元素的值和索引。对应到 C 实现TensorTopK会在函数中新增OP_TOPK算子并将TOPK_AXIS、TOPK_KVALUE、TOPK_ORDER、topkAlgo等属性写入算子属性值结果与索引结果作为两个输出同时产出见 framework/src/interface/operation/vector/topk.cpp。算法说明MERGE_SORT 与 RADIX_SELECTalgo参数的取值由枚举TopKAlgo定义完整定义见 TopKAlgo在 Python 侧通过 pybind11 导出见 python/src/bindings/enum.cppclass TopKAlgo(enum.Enum): MERGE_SORT ... # 归并排序算法 RADIX_SELECT ... # 基数选择算法参数值说明MERGE_SORT归并排序算法。对整个张量排序之后选出前k个数。RADIX_SELECT基数选择算法。先找出第k个数之后根据第k个数找出前k个数。使用建议默认行为如果不指定算法默认使用MERGE_SORT模式性能要求高的场景推荐使用RADIX_SELECT模式时间复杂度为 O(n)对性能要求不高的场景可以使用MERGE_SORT模式时间复杂度为 O(nlogn)。从源码结构看MERGE_SORT路径在 framework/src/interface/operation/vector/topk.cpp 中被展开为三阶段算子流水OP_BITSORT位排序产生带临时空间的排序中间结果→OP_TILEDMRGSORT按 4 路多队列归并TOPK_MERGE_SIZE固定为 32→OP_EXTRACT以maskMode0抽取前 k 个值、maskMode1抽取前 k 个索引而RADIX_SELECT路径则直接生成单个OP_RADIX_SELECT算子其临时空间由RadixSelectGetCalcTopKTempBlockCount6/12/22 block与RadixSelectGetSortTempBlockCount22/26/34 block两个函数按输入数据类型字节数计算见 framework/src/interface/operation/vector/topk.cpp。约束说明使用Tensor.topk/pypto.topk时须遵守以下约束只支持对尾轴进行 topk 操作选用MERGE_SORT算法时TileShape 尾轴需要小于 22KBTileShape[-1]*4 22KB选用RADIX_SELECT算法时由于存在临时内存使用设置 TileShape 时需保证输入 Tile、输出 Tile 及临时空间的总占用小于可用 UB。记 TileShape 的次尾轴为 tileH若不存在则为 1尾轴为 tileWtileW 对齐到 128 记为 tileWAlign各数据类型对应的临时空间如下输入数据类型临时空间大小字节DT_INT8、DT_UINT8、DT_INT16、DT_UINT16、DT_FP16、DT_BF16tileH * max(6 * tileWAlign max(3072, 8 * tileWAlign), 22 * tileWAlign)DT_FP32、DT_INT32、DT_UINT32tileH * max(12 * tileWAlign max(3072, 8 * tileWAlign), 26 * tileWAlign)DT_INT64、DT_UINT64tileH * max(22 * tileWAlign max(3072, 8 * tileWAlign), 34 * tileWAlign)表格中的 max(a,b) 表示取 a 和 b 中的最大值。可以看到临时空间主要由“计算 TopK 工作区6/12/22 个 tileWAlign 块 分支工作区max(3072, 8*tileWAlign)”与“排序工作区22/26/34 个 tileWAlign 块”取较大者构成与源码中calcTopKWorkspace、sortWorkspace的计算逻辑一一对应。选用RADIX_SELECT算法时尾轴不可切分TileShape[-1]必须大于等于input.shape[-1]k TileShape[-1] k input.shape[-1]RADIX_SELECT算法在不同型号的支持度Ascend 950PR/Ascend 950DT支持Atlas A3 训练系列产品/Atlas A3 推理系列产品不支持Atlas A2 训练系列产品/Atlas A2 推理系列产品不支持Tensor 类型输入不支持TileOpFormat.TILEOP_NZ格式。其中约束 4 的实现依据可直接在源码中找到TiledTopK分派到RADIX_SELECT路径时会检查operand-shape[size - 1] tileShape.GetVecTile()[size - 1]不满足即抛出The tile_shape[-1] should greater than or equal to input.shape[-1]错误见 framework/src/interface/operation/vector/topk.cpp。调用示例TileShape 设置示例说明调用该 operation 接口前应通过set_vec_tile_shapes设置 TileShape。TileShape 维度应和输入 input 一致。示例 1输入 input shape 为[m, n, p]dim 为 2largest 为 True输出为[m, n, k]TileShape 设置为[m1, n1, p1]则 m1n1p1 分别用于切分 mnp 轴。p1 必须大于等于 kk 轴不支持切分必须保证全载。pypto.set_vec_tile_shapes(4, 16, 32)示例 2输入 input shape 为[2, 3]尾轴长度为 3使用RADIX_SELECT时需保证TileShape[-1] 3例如pypto.set_vec_tile_shapes(1, 32) # 尾轴 32 对齐到 128满足尾轴不可切分约束接口调用示例x pypto.tensor([2, 3], pypto.DT_FP32) y pypto.topk(x, 2, -1, True, pypto.TopKAlgo.MERGE_SORT)结果示例如下输入数据x: [[1.0 2.0 3.0], [1.0 2.0 3.0]] 输出数据y[0]: [[3.0 2.0], [3.0 2.0]] 输出数据y[1]: [[2, 1], [2, 1]]该示例同样出现在 python/pypto/op/comparison.py 的 docstring 中可作为验证预期结果的依据每行取最大的 2 个值y[0]为值[3.0, 2.0]y[1]为对应的原始索引[2, 1]。使用 Tensor.topk 方法形式与顶层接口等价的方法形式默认使用MERGE_SORTx pypto.tensor([2, 3], pypto.DT_FP32) values, indices x.topk(2, -1, True) # dim 缺省时等价于 x.topk(2)使用 RADIX_SELECT 算法在 Ascend 950 系列产品上如需更高的性能O(n) 复杂度可显式指定RADIX_SELECT输入需为支持的数据类型例如DT_FP16、DT_INT32等x pypto.tensor([2, 8], pypto.DT_FP16) pypto.set_vec_tile_shapes(1, 32) # 尾轴须 8 且不可切分 values, indices pypto.topk(x, 3, -1, True, pypto.TopKAlgo.RADIX_SELECT)注意RADIX_SELECT仅在 Ascend 950PR/Ascend 950DT 上支持在 Atlas A2/A3 系列产品上使用会直接报错此时应回退到默认的MERGE_SORT算法仅支持DT_FP32输入。常见用法场景与注意事项归一化与注意力掩码沿特征维取 top-k 索引是稀疏注意力等场景的常用预处理步骤输出indices为DT_INT32可直接用于后续gather类索引操作动态 shape 支持源码中TensorTopK会根据输入张量的动态有效 shapeDynValidShape同步更新输出的有效 shape将dim轴的有效长度更新为k因此动态 shape 场景下可直接使用本接口内存规划当输入尾轴较长、单 Tile 无法承载时MERGE_SORT支持尾轴内部分块排序后逐级归并axisTileNum分块 4 路多队列归并而RADIX_SELECT要求尾轴全载需按约束 3 的公式预先估算临时空间保证输入 Tile、输出 Tile 与临时空间三者之和不超过可用 UB与argsort的区分仓库 python/pypto/op/comparison.py 中还提供了argsort仅返回排序索引若只需索引不需要值可结合业务选择更轻量的接口。小结Tensor.topk是 PyPTO 面向尾轴 top-k 计算的一站式接口方法形式简洁x.topk(k, dim, largest)顶层算子pypto.topk则额外暴露algo以支持算法选型。理解MERGE_SORTO(nlogn)全平台支持仅DT_FP32与RADIX_SELECTO(n)仅 Ascend 950 系列支持 11 种数据类型的差异并严格遵循 TileShape 尾轴约束与 UB 临时空间计算公式即可在 Atlas A2/A3 及 Ascend 950 系列产品上正确、高效地使用该算子。【免费下载链接】pyptoPyPTO发音: pai p-t-oParallel Tensor/Tile Operation编程范式。项目地址: https://gitcode.com/cann/pypto创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
网站建设高端定制企业官网