PyPTO 逐元素大于比较算子 pypto.gt:函数原型、广播约束与 TileShape 切分实践
【免费下载链接】pyptoPyPTO(发音: pai p-t-o):Parallel Tensor/Tile Operation编程范式。项目地址: https://gitcode.com/cann/pypto
导读:
pypto.gt是 CANN PyPTO(Parallel Tensor/Tile Operation 编程范式)提供的逐元素大于比较算子,用于判断输入 Tensor 各位置元素是否严格大于另一个 Tensor 或标量,返回同 Shape 的 DT_BOOL 结果。本文以仓库文档 docs/zh/api/tensor_api/operation/pypto-gt.md 为核心,结合源码实现与测试用例,完整讲解pypto.gt的函数原型、参数与返回语义、各型号产品的数据类型约束、广播与 TileShape 设置要求,并给出可直接复用的调用示例与实战验证方法。
一、功能定位与适用场景
pypto.gt实现的是逐元素大于比较(element-wise greater-than)运算,即对输入 Tensor 中每个位置逐一执行input[i] > other[i]的比较,比较结果以布尔值(True/False)逐位写出到输出 Tensor。它是 PyPTO 张量运算体系中比较类算子的核心成员,与pypto.ge(大于等于)、pypto.eq(等于)、pypto.ne(不等于)、pypto.lt(小于)、pypto.le(小于等于)共同构成完整的六元比较算子族。
该算子常见的应用场景包括:
- 掩码(mask)生成:如激活函数、归一化、注意力等算子中,依据数据是否超过阈值生成布尔掩码;
- 数值截断与条件选择:与
pypto.where等算子配合,实现基于比较结果的条件分支数据流; - 统计与过滤:统计满足某一阈值的元素个数、位置等。
从实现上看,pypto.gt是pypto.greater的别名,二者在 python/pypto/op/comparison.py 中为完全等价的实现,内部统一调用底层 C++ 算子npu::tile_fwk::Compare,比较模式取OpType::GT、输出类型取OutType::BOOL。
二、产品支持情况
pypto.gt在以下产品系列上均得到支持(依据 pypto-gt.md):
- Ascend 950PR / Ascend 950DT:支持;
- Atlas A3 训练系列产品 / Atlas A3 推理系列产品:支持;
- Atlas A2 训练系列产品 / Atlas A2 推理系列产品:支持。
不同产品系列对input/other所允许的数据类型存在差异,详见下文"约束说明"一节的逐型号数据类型表。
三、函数原型
gt(input: Tensor, other: Union[Tensor, float, Element]) -> Tensor其中:
input:第一个源操作数,类型为 Tensor;other:第二个源操作数,支持 Tensor、float 或 Element 三种类型;- 返回值为 Tensor 类型(DT_BOOL)。
与原型完全对应的实现位于 python/pypto/op/comparison.py:
@op_wrapper def gt(input: Tensor, other: Union[Tensor, float, Element]) -> Tensor: if isinstance(other, float): # Tensor vs Scalar comparison return pypto_impl.Compare(input, pypto_impl.Element(input.dtype, other), OpType.GT, OutType.BOOL) return pypto_impl.Compare(input, other, OpType.GT, OutType.BOOL)从这里可以读出两条关键实现细节:
- float 自动转 Element:当
other传入 Python 的float时,实现会以input的 dtype 构造pypto_impl.Element(input.dtype, other),即该浮点标量默认按DT_FP32语义参与比较;如需使用其他数据类型(如 DT_FP16、DT_INT16 等),则应显式通过Element构建(见 python/pypto/_element.py,Element(dtype, data)支持 int/float 标量)。 - 底层统一走 Compare 算子:不论
other是 Tensor 还是标量,最终都落入 C++ 层的Compare。对应声明见 framework/include/tilefwk/tilefwk_op.h,pybind 绑定见 python/src/bindings/operation.cpp,其中比较模式枚举OpType包含EQ / NE / LT / LE / GT / GE,输出类型枚举OutType包含BOOL / BIT(见 tilefwk_op.h),pypto.gt固定使用GT + BOOL组合。
四、参数说明
| 参数名 | 输入/输出 | 说明 |
|---|---|---|
| input | 输入 | 源操作数。支持的类型为 Tensor。不同型号支持的 Tensor 数据类型有所差异,详细请参见约束说明。不支持空 Tensor;Shape 仅支持 1-4 维;Shape Size 不大于 2147483647(即 INT32_MAX)。 |
| other | 输入 | 源操作数。支持的类型为 Tensor、float、Element。当为 float 类型时会自动转换为 Element 类型,float 对应 DT_FP32;当需要使用其他数据类型时,可以通过 Element 构建。不同型号支持的 Tensor 和 Element 的数据类型有所差异,详细请参见约束说明。不支持空 Tensor;Shape 仅支持 1-4 维;Shape Size 不大于 2147483647(即 INT32_MAX)。 |
补充说明:
input与other的数据类型须保持一致(见下文约束第 1 条);other为 Tensor 时,其 Shape 支持与input相同或可通过广播对齐(见下文约束第 2 条);- 维度上限为 4 维,且元素总数不能超过 INT32_MAX,这是由 Tile 框架内部索引与缓冲区寻址的 32 位粒度决定的,在构造大规模输入时需要提前核算。
五、返回值说明
返回一个Shape 与输入 Tensor 一致、数据类型为 DT_BOOL 的 Tensor:
- 若
input对应位置的元素值严格大于other对应位置的元素值,则返回 True; - 其余位置返回 False。
注意"严格大于"的语义:相等或小于均不满足条件,这正是gt与ge(greater-or-equal)的关键区别。比较输出统一为布尔张量,便于后续与逻辑算子、选择算子衔接。
六、约束说明
- 类型一致性:
input和other类型须保持一致(other为 float 时按其自动转换为的 Element 语义理解)。 - 广播支持:支持多维度广播到相同形状,即
other的 Shape 可以比input更小,按 PyPTO 的广播规则扩展后再逐元素比较。 - Tensor 和 Element 数据类型,按产品型号区分:
| 产品系列 | 支持的数据类型 |
|---|---|
| Ascend 950PR / Ascend 950DT | DT_FP16、DT_FP32、DT_INT16、DT_INT64、DT_UINT64 |
| Atlas A3 训练系列产品 / Atlas A3 推理系列产品 | DT_FP16、DT_FP32 |
| Atlas A2 训练系列产品 / Atlas A2 推理系列产品 | DT_FP16、DT_FP32 |
可以看到,Ascend 950 系列额外支持整型比较(DT_INT16 / DT_INT64 / DT_UINT64),而 A2 / A3 系列当前仅支持浮点类型 DT_FP16 / DT_FP32。据此,若在 A2/A3 产品上对整型数据执行pypto.gt,需要先进行类型转换。 4.格式限制:Tensor 类型输入不支持TileOpFormat.TILEOP_NZ格式,即参与比较的 Tensor 需使用 PyPTO 默认的矢量数据排布格式。
七、调用示例
7.1 TileShape 设置示例
PyPTO 的矢量(vector)类算子在执行前需要通过pypto.set_vec_tile_shapes设置 TileShape(分块形状),该接口的定义见 docs/zh/api/tensor_api/config/pypto-set_vec_tile_shapes.md,原型为set_vec_tile_shapes(*args: int) -> None,最多可传 4 个维度,且每个维度必须大于 0。
调用pypto.gt前,TileShape 维度应和输出一致,并按输出张量的各轴进行切分。
示例 1:非广播场景
输入input的 Shape 为[m, n],other的 Shape 为[m, n],输出 Shape 为[m, n]。此时 TileShape 设置为[m1, n1],m1、n1分别用于切分 m 轴、n 轴:
pypto.set_vec_tile_shapes(4, 16)示例 2:广播场景
输入input的 Shape 为[m, n],other的 Shape 为[m, 1],输出 Shape 为[m, n]。TileShape 仍设置为[m1, n1],m1、n1分别用于切分 m 轴、n 轴:
pypto.set_vec_tile_shapes(4, 16)即广播发生在比较语义层面([m, 1]沿 n 轴扩展为[m, n]),TileShape 依然按最终输出 Shape 的两个维度来设置。
7.2 接口调用示例
最简单的 Tensor-Tensor 比较:
a = pypto.tensor([3], pypto.DT_FP32) b = pypto.tensor([3], pypto.DT_FP32) out = pypto.gt(a, b)结果示例如下:
输入数据a: [1.0 2.0 3.0] 输入数据b: [2.0 2.0 2.0] 输出数据out: [False, False, True]即仅当a的元素严格大于b的对应元素(3.0 > 2.0)时为 True,1.0 > 2.0、2.0 > 2.0均为 False。
标量比较(other传 float)等价写法:
a = pypto.tensor([3], pypto.DT_FP32) out = pypto.gt(a, 2.0) # 等价于 pypto.gt(a, pypto.Element(pypto.DT_FP32, 2.0))7.3 完整 Kernel 级实战:分块循环 + TileShape + 结果回写
仅调用pypto.gt会生成一个比较算子,但在 PyPTO 中通常需要配合pypto.function、pypto.loop、pypto.view、pypto.assemble构成完整的 Kernel。仓库测试 python/tests/st/operation/vector/test_greater.py 提供了可复用的完整范式,下面以test_vector_operation_greater为例说明关键环节:
import numpy as np import torch import pypto shape = (32, 32) view_shape = (16, 16) tile_shape = (8, 8) a = pypto.tensor(shape, pypto.DT_FP32, "Greater_TENSOR_a") b = pypto.tensor(shape, pypto.DT_FP32, "Greater_TENSOR_b") c = pypto.tensor(shape, pypto.DT_BOOL, "Greater_TENSOR_c") with pypto.function("Greater", a, b, c): for b_idx in pypto.loop(2, name="LOOP_GREATER_L0", idx_name="b_idx"): for s_idx in pypto.loop(2, name="LOOP_GREATER_L1", idx_name="s_idx"): tile_a = pypto.view(a, view_shape, [b_idx * view_shape[0], s_idx * view_shape[1]]) tile_b = pypto.view(b, view_shape, [b_idx * view_shape[0], s_idx * view_shape[1]]) pypto.set_vec_tile_shapes(tile_shape[0], tile_shape[1]) tile_a.move(pypto.greater(tile_a, tile_b)) # gt 与 greater 等价 pypto.assemble(tile_a, [b_idx * view_shape[0], s_idx * view_shape[1]], c)该示例的核心链路是:
- 用
pypto.view从全局张量上切出view_shape大小的分块(tile); - 调用
pypto.set_vec_tile_shapes(tile_shape[0], tile_shape[1])设置矢量算子的内部 TileShape; - 在分块上执行
pypto.greater(tile_a, tile_b)(与pypto.gt完全等价),结果move回分块; - 用
pypto.assemble将比较结果写回全局输出张量c。
同文件的test_greater_scalar(test_greater.py)则演示了other传标量50.0的写法,即view_block.move(pypto.greater(view_block, scalar_value))。
验证方式:测试中使用torch.npu.set_device(device_id)指定设备,运行后与torch.greater(a_tensor, b_tensor)(标量场景为torch.greater(input_data, scalar_value))的结果做assert_allclose对比,用于校验设备侧比较结果与参考实现一致。
八、与比较算子族的横向对比
pypto.gt位于 PyPTO 比较算子族(python/pypto/op/comparison.py)中,与同族算子的语义差异如下(均为逐元素、输出 DT_BOOL):
| 算子 | 语义 | 示例(a=[1,2,3], b=[2,2,2]) |
|---|---|---|
pypto.gt(a, b) | 严格大于 | [False, False, True] |
pypto.ge(a, b) | 大于等于 | [False, True, True] |
pypto.eq(a, b) | 等于 | [False, True, False] |
pypto.ne(a, b) | 不等于 | [True, False, True] |
pypto.lt(a, b) | 严格小于 | [True, False, False] |
pypto.le(a, b) | 小于等于 | [True, True, False] |
从源码看,ge/eq/ne/lt/le与gt的实现结构完全一致,仅在OpType枚举取值上不同(见 comparison.py),因此本文关于广播、TileShape、数据类型约束的结论对整族算子同样适用。选择哪个算子只需根据比较语义(严格 / 非严格)确定。
九、常见问题与使用建议
- 数据类型不匹配:
input与other数据类型须一致,标量传入 float 时自动按 DT_FP32 处理。若要与其他类型比较,应显式Element(dtype, value)构造,避免隐式类型歧义。 - 型号不支持整型比较:A2 / A3 系列仅支持 DT_FP16 / DT_FP32,若需要对整型数据比较,可先通过
pypto.cast转换类型后再调用pypto.gt(详见 docs/zh/api/tensor_api/operation/pypto-cast.md)。 - 忘记设置 TileShape:矢量算子执行前须调用
pypto.set_vec_tile_shapes,且 TileShape 各维度需大于 0、维度数与输出一致;否则算子无法正确切分数据。 - NZ 格式限制:参与比较的 Tensor 不要使用
TileOpFormat.TILEOP_NZ格式排布。 - Shape 上限:输入 Shape 仅支持 1-4 维,且 Shape Size 不大于 INT32_MAX,超大张量需先拆分处理。
十、总结
pypto.gt是 PyPTO 中实现逐元素"严格大于"比较的标准算子:功能定义在 docs/zh/api/tensor_api/operation/pypto-gt.md,Python 层入口为 python/pypto/op/comparison.py,底层由 C++Compare(GT, BOOL)算子承载(framework/include/tilefwk/tilefwk_op.h、python/src/bindings/operation.cpp),并有完整的 Kernel 级测试用例 python/tests/st/operation/vector/test_greater.py 佐证。掌握其"严格大于、输出 BOOL、支持广播、需设 TileShape"四大要点,即可在算子开发中正确、高效地使用比较运算,并可类推至同族六个比较算子。
【免费下载链接】pyptoPyPTO(发音: pai p-t-o):Parallel Tensor/Tile Operation编程范式。项目地址: https://gitcode.com/cann/pypto
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考