PyPTO vf.gt 逐元素大于比较算子:SIMD 向量编程中的掩码比较与条件选择实战
2026/9/21 15:13:07 网站建设 项目流程
  • 人工智能
  • 编译器
  • 模型编译
  • 高性能计算
  • 深度学习
  • CANN

【免费下载链接】pypto

PyPTO(发音: pai p-t-o):Parallel Tensor/Tile Operation编程范式。

项目地址:https://gitcode.com/cann/pypto
点击查看免费下载

导读

vf.gt是 PyPTO 向量函数(vector_function)编程体系中用于逐元素大于比较的核心 SIMD 算子,它以寄存器级粒度比较两个源操作数,并把比较结果写入掩码寄存器(mask_reg)的对应比特位,是构建max、条件赋值、数据过滤等向量运算的基础构件。本文以 gt.md 为骨架,结合 PyPTO 仓库源码与其配套的 reg_tensor、mask_reg 文档,完整讲解该算子的产品支持情况、语义、参数约束、返回机制,并给出可直接运行的 FP32 与 INT64 调用示例。读完本文,你将掌握如何用vf.gt生成比较掩码,并配合vf.select完成寄存器级的条件选择流水。

产品支持情况

vf.gt与 PyPTO 向量寄存器体系(reg_tensor / mask_reg)的产品支持范围一致:

  • Ascend 950PR / Ascend 950DT:支持
  • Atlas A3 训练系列产品 / Atlas A3 推理系列产品:不支持
  • Atlas A2 训练系列产品 / Atlas A2 推理系列产品:不支持

使用前请确认目标 NPU 架构,当前仓库文档中该算子的硬件指令路径仅面向 Ascend 950 系列产品实现。

功能说明

vf.gtsrc0src1的每个元素执行大于(>)比较,将比较结果逐比特写入目的操作数dst_mask中对应数据元素的掩码位:

  • src0_i > src1_i为真,则dst_mask中该元素对应的比特位为1
  • 否则该比特位为0

其逐元素语义可用如下公式描述:

$$dstReg_i = \begin{cases} 1 & \text{if } src0_i > src1_i \ 0 & \text{otherwise} \end{cases}$$

两个值得注意的特性:

  1. 标量/向量自动分发:第二个参数src1可以是标量,也可以是 reg_tensor。接口会自动识别参数形态并分发到对应的硬件指令路径(vector-scalar 比较路径或 vector-vector 比较路径)。这一点在源码 python/pypto_pro/language/_vf_api.py 的gt声明注释中有明确体现:"If the second argument is a scalar literal the vector-scalar compare path is used; otherwise the vector-vector compare path is used."
  2. 掩码驱动:比较结果不是普通数据寄存器,而是掩码寄存器dst_mask,后续可被vf.select等消费掩码的算子直接使用,实现"比较 → 选择"的组合流水。

函数原型

gt(src0, src1, preg) -> dst_mask
  • 返回值类型:dst_mask(mask_reg,存放比较结果)。
  • 函数为vf向量指令空间内的静态方法,在@pl.vector_function修饰的向量函数内调用。

参数说明

参数输入/输出说明
src0输入源操作数,reg_tensor。支持的数据类型为:DT_INT8、DT_UINT8、DT_INT16、DT_UINT16、DT_FP16、DT_BF16、DT_INT32、DT_UINT32、DT_FP32、DT_INT64、DT_UINT64。src0 和 src1 可以是同一个 reg_tensor。
src1输入比较操作数,可以是标量或 reg_tensor,数据类型与 src0 一致。
preg输入mask_reg,指定参与比较的元素范围。通过 preg 参数控制的未选中元素在目的操作数中被置零。

参数细节补充

  • 数据形态与元素粒度:reg_tensor 总大小固定为 256 字节,元素个数由 dtype 决定(如 FP32 为 64 个元素、INT64 为 32 个元素),详见 reg_tensor.md 中的数据类型约束表。
  • 掩码粒度与位宽:mask_reg 总位宽固定为 256 bit,其粒度由 dtype 决定——例如 b32 粒度(FP32/INT32/UINT32)下 64 个元素共占 256 bit,b64 粒度(INT64/UINT64)下 32 个元素占 256 bit。因此vf.gt的结果掩码位数与参与比较的数据类型一一对应,详见 mask_reg.md。
  • preg 的过滤语义:preg 中比特位为 0(无效)的元素不参与运算,且目的掩码对应位置置零;只有比特位为 1(有效)的元素参与比较并写入结果。

约束说明

无。

返回值说明

返回dst_mask,类型为目标 mask_reg,存放逐元素比较结果。该掩码寄存器由编译器在赋值形式中自动声明(如gt_mask = vf.gt(reg_a, reg_b, cmp_mask)),在 vector_function 函数内创建和使用,函数结束后自动释放。

调用示例

基本调用示例(FP32)

以下示例演示完整的"加载 → 比较 → 条件选择 → 存储"流程:用vf.gt生成a > b的掩码,再通过vf.select按掩码从两个源寄存器中挑选较大值,等价于torch.where(a > b, a, b)(即逐元素 max):

import os import pypto_pro.language as pl import torch import torch_npu @pl.vector_function def example_vf(src_a, src_b, dst_tile): preg = vf.create_mask(pattern=pl.MaskPattern.ALL, dtype=pl.DT_FP32) reg_a = vf.load_align(src_a, 0) reg_b = vf.load_align(src_b, 0) dst_mask = vf.gt(reg_a, reg_b, preg) reg_out = vf.select(reg_a, reg_b, dst_mask) vf.store_align(dst_tile, reg_out, preg) @pl.jit() def example_kernel( a: pl.Tensor[[pl.DYNAMIC, pl.DYNAMIC], pl.DT_FP32], b: pl.Tensor[[pl.DYNAMIC, pl.DYNAMIC], pl.DT_FP32], out: pl.Tensor[[pl.DYNAMIC, pl.DYNAMIC], pl.DT_FP32], ): tf = pl.TileType(shape=[1, 64], dtype=pl.DT_FP32, target_memory=pl.MemorySpace.Vec) in_a_grp = pl.make_tile_group(type=tf, addrs=0x0, mutex_ids=[0]) in_a = in_a_grp.current() in_b_grp = pl.make_tile_group(type=tf, addrs=0x100, mutex_ids=[1]) in_b = in_b_grp.current() t_out_grp = pl.make_tile_group(type=tf, addrs=0x200, mutex_ids=[2]) t_out = t_out_grp.current() with pl.section_vector(): pl.load(in_a, a, [0, 0]) pl.load(in_b, b, [0, 0]) example_vf(in_a, in_b, t_out) pl.store(out, t_out, [0, 0]) def test_example(): device_id = int(os.environ.get("TILE_FWK_DEVICE_ID", 0)) device = f"npu:{device_id}" core_nums = 1 torch.npu.set_device(device) a = torch.randn([1, 64], device=device, dtype=torch.float32) b = torch.randn([1, 64], device=device, dtype=torch.float32) out = torch.empty([1, 64], device=device, dtype=torch.float32) example_kernelNone, core_nums torch.npu.synchronize() torch.testing.assert_close(out, torch.where(a > b, a, b), rtol=1e-5, atol=1e-5) if __name__ == "__main__": test_example() print("PASSED")

要点拆解:

  • vf.create_mask(pattern=pl.MaskPattern.ALL, dtype=pl.DT_FP32)创建全有效掩码 preg,声明在 python/pypto_pro/language/_vf_api.py 附近,其 dtype 需与寄存器数据类型一致。
  • vf.load_align/vf.store_align负责 UB Tile 与寄存器之间的数据搬运。
  • vf.gt产生掩码后,vf.select(src0, src1, dst_mask)按掩码逐元素选择(掩码位为 1 时取 src0,否则取 src1),其实现声明位于同一文件的select定义处(python/pypto_pro/language/_vf_api.py)。
  • 端到端正确性由torch.testing.assert_close(out, torch.where(a > b, a, b), ...)校验,数值与语义完全对齐。

INT64 数据类型示例

vf.gt支持 64 位整型比较。注意在 b64 粒度下,寄存器元素个数为 32,且示例中将"源谓词"与"比较结果掩码"分开声明:

import os import pypto_pro.language as pl import torch import torch_npu @pl.vector_function def example_vf_int64(src_a, src_b, dst_tile): preg = vf.create_mask(pattern=pl.MaskPattern.ALL, dtype=pl.DT_INT64) cmp_mask = vf.create_mask(pattern=pl.MaskPattern.ALL, dtype=pl.DT_INT64) reg_a = vf.load_align(src_a, 0) reg_b = vf.load_align(src_b, 0) gt_mask = vf.gt(reg_a, reg_b, cmp_mask) reg_out = vf.select(reg_a, reg_b, gt_mask) vf.store_align(dst_tile, reg_out, preg) @pl.jit() def example_kernel_int64( a: pl.Tensor[[pl.DYNAMIC, pl.DYNAMIC], pl.DT_INT64], b: pl.Tensor[[pl.DYNAMIC, pl.DYNAMIC], pl.DT_INT64], out: pl.Tensor[[pl.DYNAMIC, pl.DYNAMIC], pl.DT_INT64], ): tf = pl.TileType(shape=[1, 32], dtype=pl.DT_INT64, target_memory=pl.MemorySpace.Vec) in_a_grp = pl.make_tile_group(type=tf, addrs=0, mutex_ids=[0]) in_a = in_a_grp.current() in_b_grp = pl.make_tile_group(type=tf, addrs=256, mutex_ids=[1]) in_b = in_b_grp.current() t_out_grp = pl.make_tile_group(type=tf, addrs=512, mutex_ids=[2]) t_out = t_out_grp.current() with pl.section_vector(): pl.load(in_a, a, [0, 0]) pl.load(in_b, b, [0, 0]) example_vf_int64(in_a, in_b, t_out) pl.store(out, t_out, [0, 0]) def test_example_int64(): device_id = int(os.environ.get("TILE_FWK_DEVICE_ID", 0)) device = f"npu:{device_id}" core_nums = 1 torch.npu.set_device(device) a = torch.randint(-100, 100, [1, 32], device=device, dtype=torch.int64) b = torch.randint(-100, 100, [1, 32], device=device, dtype=torch.int64) out = torch.empty([1, 32], device=device, dtype=torch.int64) example_kernel_int64None, core_nums torch.npu.synchronize() torch.testing.assert_close(out, torch.where(a > b, a, b), rtol=0, atol=0) if __name__ == "__main__": test_example_int64() print("PASSED")

INT64 示例的差异点:

  • Tile 形状为[1, 32],与 b64 粒度下每寄存器 32 个元素对应(参考 reg_tensor.md 的数据类型约束表)。
  • 输入使用torch.randint(-100, 100, ...)生成整数数据,校验时rtol=0, atol=0精确比对整型结果。
  • 结果仍等价于torch.where(a > b, a, b),即取逐元素较大值。

与其他比较算子的配套使用

vf.gt属于 comparison_and_selection 比较与选择算子族,同族算子还包括vf.eq(等于)、vf.ne(不等于)、vf.lt(小于)、vf.le(小于等于)、vf.ge(大于等于)以及消费掩码的vf.selectvf.squeeze。这些算子共享同一种"比较产出掩码、掩码驱动选择/压缩"的编程范式,源码中全部以静态方法声明于 python/pypto_pro/language/_vf_api.py,接口签名与vf.gt保持严格一致((src0, src1, preg) -> mask),掌握vf.gt即可触类旁通地使用整个比较算子族。通过组合vf.gt+vf.select,可高效实现maxmin、条件三元选择、元素过滤等常见的向量化逻辑。

  • 人工智能
  • 编译器
  • 模型编译
  • 高性能计算
  • 深度学习
  • CANN

【免费下载链接】pypto

PyPTO(发音: pai p-t-o):Parallel Tensor/Tile Operation编程范式。

项目地址:https://gitcode.com/cann/pypto
点击查看免费下载

相关推荐

上一篇:重复文件清理指南:Czkawka 与 Krokiet 完整上手
下一篇:JMESPath PHP 项目教程

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询