CANN PyPTO SIMD 逻辑计算 API 详解:vf.and_ / or_ / xor / not_ / shift_left / shift_right
2026/9/20 3:08:45 网站建设 项目流程
  • 人工智能
  • 编译器
  • 模型编译
  • 高性能计算
  • 深度学习
  • CANN

【免费下载链接】pypto

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

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

逻辑计算(Logical Computation)是 CANN PyPTO 向量函数(Vector Function,vf)寄存器计算族中用于按位操作的核心集合。本文基于 docs/zh/api/pro_api/SIMD-API/reg_computation/logical_computation/index.md 及其六个子文档,系统讲解vf.and_vf.or_vf.xorvf.not_vf.shift_leftvf.shift_right六个按位运算接口的功能语义、函数原型、参数与数据类型约束、mask 谓词行为,并结合完整可运行的 Kernel 示例与仓库源码,帮助开发者在 Ascend 950 系列产品的 Vector 单元上编写正确的位级运算代码。

一、接口总览与适用场景

逻辑计算接口属于 Reg 计算 目录下的一个独立子类,主要解决向量寄存器级别的位级数据处理需求。该族接口共 6 个:

接口运算函数原型
vf.and_按位与and_(src0, src1, preg, mode) -> dst
vf.or_按位或or_(src0, src1, preg, mode) -> dst
vf.xor按位异或xor(src0, src1, preg, mode) -> dst
vf.not_按位取反not_(src, preg, mode) -> dst
vf.shift_left左移shift_left(src, shift, preg, mode) -> dst
vf.shift_right右移shift_right(src, shift, preg, mode) -> dst

典型应用场景包括:

  • 掩码(mask)逻辑组合:由比较指令(如vf.gevf.lt)产出的mask_reg可以通过and_/or_/xor/not_进行布尔组合,构造复杂谓词后再驱动后续向量操作。
  • 位域提取与打包:通过shift_left/shift_right配合掩码完成定点数位域抽取、量化位宽的搬移等操作。
  • 无分支条件计算:将条件结果编码为掩码/位模式,通过按位运算实现选择与归一化,避免分支跳转。

产品支持情况

6 个接口的产品支持情况完全一致,均在Ascend 950PR / Ascend 950DT 上支持,在 Atlas A3 训练/推理系列产品与 Atlas A2 训练/推理系列产品上不支持。编写算子时请先确认目标硬件平台,否则会因指令不可用导致编译或运行失败。

二、公共语义:preg 谓词与 MergeMode

除移位接口的shift参数外,所有逻辑计算接口共享同一套参数与返回语义:

  • src / src0 / src1(输入):源操作数,类型为 reg_tensor 或 mask_reg;双操作数接口要求src0src1与目的操作数dst的数据类型保持一致。
  • preg(输入):谓词掩码寄存器 mask_reg,用于逐 lane 筛选哪些元素参与运算。
  • mode(输入,可选):对应 MergeMode 枚举。当前仅支持默认值pypto_pro.language.MergeMode.ZEROING,即preg 未筛选(无效)的元素在 dst 中直接置 0MergeMode.MERGING当前不支持。
  • 返回值(dst):目的操作数,reg_tensormask_reg类型,支持的数据类型与源操作数一致。

需要注意的是,vf.xor的 源码实现 中额外暴露了一个可选dtype参数(如pl.DT_UINT16),用于类型特化变体;文档公开原型为xor(src0, src1, preg, mode),常规调用按文档原型传入即可。

从源码角度看,这些接口均以@staticmethod+@_api_decl方式声明在vf命名空间中,文档注释中明确写明了逐 lane 语义:对每个mask[i]为激活态的 lanei计算并写入dst[i](见 python/pypto_pro/language/_vf_api.py)。

三、按位与 / 或 / 异或:and_、or_、xor

3.1 功能语义

三个双操作数接口按位逐元素运算:

  • vf.and_dstReg_i = srcReg0_i & srcReg1_i
  • vf.or_dstReg_i = srcReg0_i | srcReg1_i
  • vf.xordstReg_i = srcReg0_i ^ srcReg1_i

支持的数据类型(三者一致)为:DT_INT8DT_UINT8DT_INT16DT_UINT16DT_FP16DT_BF16DT_INT32DT_UINT32DT_FP32DT_INT64DT_UINT64DT_FP8E4M3FNDT_FP8E5M2DT_FP8E8M0。注意该集合不仅覆盖整型,也覆盖浮点与 8bit 浮点格式——按位运算不关心数值解释,直接作用于二进制位模式。vf.not_支持的位宽范围略窄:DT_INT8DT_UINT8DT_INT16DT_UINT16DT_INT32DT_UINT32DT_FP16DT_FP32DT_INT64DT_UINT64(不含 BF16 与 FP8 系列)。

3.2 reg_tensor 调用示例(and_)

以下示例完整演示"加载 → 按位与 → 存储"的向量函数 + Kernel 全流程(基于 and_.md):

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_UINT16) reg_a = vf.load_align(src_a, 0) reg_b = vf.load_align(src_b, 0) reg_out = vf.and_(reg_a, reg_b, preg) vf.store_align(dst_tile, reg_out, preg) @pl.jit() def example_kernel( a: pl.Tensor[[pl.DYNAMIC, pl.DYNAMIC], pl.DT_UINT16], b: pl.Tensor[[pl.DYNAMIC, pl.DYNAMIC], pl.DT_UINT16], out: pl.Tensor[[pl.DYNAMIC, pl.DYNAMIC], pl.DT_UINT16], ): tf = pl.TileType(shape=[1, 128], dtype=pl.DT_UINT16, 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.randint(0, 256, [1, 128], device=device, dtype=torch.int16) b = torch.randint(0, 256, [1, 128], device=device, dtype=torch.int16) out = torch.empty([1, 128], device=device, dtype=torch.int16) example_kernelNone, core_nums torch.npu.synchronize() assert out.dtype == torch.int16 if __name__ == "__main__": test_example() print("PASSED")

代码要点:

  • pl.vector_function装饰器声明向量函数体(运行在 Vector 单元上),pl.jit()声明可编译 Kernel;
  • pl.TileType(shape=[1, 128], ..., target_memory=pl.MemorySpace.Vec)将数据 tile 放到 Vector 内存;
  • pl.make_tile_group(type=tf, addrs=0x0, mutex_ids=[0])以物理地址 + 互斥 ID 声明寄存器资源,地址分别为0x00x1000x200(相邻 tile 以 256 字节间隔排布);
  • pl.section_vector()划定向量指令执行区间,内部依次完成load、向量函数、store
  • vf.load_align/vf.store_align使用对齐访问加载/回写寄存器数据。

3.3 mask_reg 调用示例:用按位运算组合掩码

当源操作数为mask_reg时,三个接口对掩码执行按位运算。这是构造复合谓词最常用的模式。以vf.xor为例,先用比较指令生成mask_a(元素 ≥ 0 为真),再与全 1 掩码异或得到"元素 < 0"的掩码,最后驱动vf.abs实现"负数取绝对值、正数清零":

@pl.vector_function def example_vf(src_tile, dst_tile): preg = vf.create_mask(pattern=pl.MaskPattern.ALL, dtype=pl.DT_FP32) reg = vf.load_align(src_tile, 0) mask_a = vf.ge(reg, 0.0, preg) mask_full = vf.create_mask(pattern=pl.MaskPattern.ALL, dtype=pl.DT_FP32) preg_xor = vf.xor(mask_a, mask_full, preg) reg_dst = vf.abs(reg, preg_xor) vf.store_align(dst_tile, reg_dst, preg)

对应vf.and_vf.or_的掩码示例同样出现在各自文档中:

  • and_preg_and = vf.and_(mask_a, mask_full, preg)后接vf.abs,语义为"非负元素取原值(abs 不变),负元素清零",Host 端用torch.where(a >= 0, a, torch.zeros_like(a))校验。
  • or_preg_or = vf.or_(mask_a, mask_b, preg),其中mask_a = vf.ge(reg, 0.0, preg)mask_b = vf.lt(reg, 0.0, preg),二者或运算后恒为全真,最终输出即torch.abs(a)
  • xorpreg_xor = vf.xor(mask_a, mask_full, preg)实现取反效果,输出为"负数取绝对值、正数清零",即torch.where(a < 0, torch.abs(a), torch.zeros_like(a))

vf.not_的掩码用法完全一致:preg_not = vf.not_(mask_a, preg),同样是"负数取绝对值、正数清零",Host 端期望torch.where(a < 0, torch.abs(a), torch.zeros_like(a))

可以看出:mask_regreg_tensor共用同一套按位运算语义,这为掩码级逻辑(如交集、并集、差集、取反)提供了与数据运算一致的编程体验。

3.4 INT64 位宽示例

三个双操作数接口与not_的文档均提供了 INT64 场景示例。INT64 元素宽度为 64 bit,Vector 寄存器容纳元素个数相应减半(示例中 tile 形状为[1, 32]),寄存器地址间隔仍为 256 字节:

@pl.vector_function def example_vf_int64(src_tile, dst_tile): preg = vf.create_mask(pattern=pl.MaskPattern.ALL, dtype=pl.DT_INT64) reg_a = vf.load_align(src_tile, 0) reg_out = vf.xor(reg_a, reg_a, preg) # a ^ a == 0 vf.store_align(dst_tile, reg_out, preg)

Host 端校验分别使用torch.testing.assert_close(out, a, ...)(and_/or_ 用a & aa | a)、out == ~a(not_)、out == 0(xor)。将同一操作数同时作为 src0/src1 的写法,可用于验证接口的原地/自运算行为并产出恒等值或全 0。

四、移位运算:shift_left 与 shift_right

移位接口与前四个接口最大的不同在于移位量 shift 的两种形态,接口会根据 shift 参数的类型自动选择模式:

  • 标量模式:shift 为整数值或标量变量,所有元素统一移动相同位数;
  • reg_tensor 模式:shift 为 reg_tensor,每个元素按对应 lane 的位数分别移动。

4.1 左移语义:逻辑左移与算术左移

vf.shift_left执行dst_i = src_i << shift_i,按源数据类型分两种行为:

  • 无符号类型 → 逻辑左移:最高位丢弃、最低位补 0。例如 DT_UINT16 的1010101010101010左移 1 位得到0101010101010100
  • 有符号类型 → 算术左移:位模式变化与逻辑左移相同(丢弃高位、低位补 0),区别在于结果按有符号类型解释。例如 DT_INT16 的1010101010101010左移 1 位位模式为0101010101010100,左移 3 位为0101010101010000

4.2 右移语义:逻辑右移与算术右移

vf.shift_right执行dst_i = src_i >> shift_i

  • 无符号类型 → 逻辑右移:最低位丢弃、最高位补 0。例如 DT_UINT16 的1010101010101010右移 1 位得到0101010101010101
  • 有符号类型 → 算术右移:最低位丢弃、最高位复制符号位。例如 DT_INT16 的1010101010101010(符号位为 1)算术右移 1 位得到1101010101010101,右移 3 位得到1111010101010101

4.3 位移量边界行为(重要)

文档对移位量超出位宽的边界情况给出了明确约定:

  • 左移(reg_tensor 模式):无论逻辑左移(无符号)还是算术左移(有符号),位移量大于数据类型位宽时输出 0
  • 右移(reg_tensor 模式):逻辑右移(无符号)位移量大于位宽输出 0;算术右移(有符号)时,src 小于 0 且位移量大于位宽输出 -1(符号扩展填满),src 大于等于 0 输出 0
  • 两种模式均不支持负数移位量,传入负数行为未定义,编码时应自行保证。

4.4 数据类型约束

移位接口的数据类型约束比按位接口严格,仅支持 8 种整型,且shift在 reg_tensor 模式下恒为有符号整型(对应位宽与 src 相同的 INT 类型):

dstsrcshift(标量模式 / reg_tensor 模式)
DT_INT8DT_INT8整型标量 / DT_INT8
DT_UINT8DT_UINT8整型标量 / DT_INT8
DT_INT16DT_INT16整型标量 / DT_INT16
DT_UINT16DT_UINT16整型标量 / DT_INT16
DT_INT32DT_INT32整型标量 / DT_INT32
DT_UINT32DT_UINT32整型标量 / DT_INT32
DT_INT64DT_INT64整型标量 / DT_INT64
DT_UINT64DT_UINT64整型标量 / DT_INT64

返回值dstreg_tensor,数据类型与 src 一致(同样受上表约束)。

4.5 标量模式与 reg_tensor 模式调用示例

标量模式(所有元素统一左移 4 位):

@pl.vector_function def example_vf_scalar(src_tile, dst_tile): preg = vf.create_mask(pattern=pl.MaskPattern.ALL, dtype=pl.DT_UINT32) reg_src = vf.load_align(src_tile, 0) reg_out = vf.shift_left(reg_src, 4, preg) vf.store_align(dst_tile, reg_out, preg)

Host 端期望out == a << 4;右移标量模式示例使用vf.shift_right(reg_src, 24, preg),Host 端期望out == a >> 24

reg_tensor 模式(逐元素移位,shift 也需先加载为寄存器):

@pl.vector_function def example_vf_vector(src_tile, shift_tile, dst_tile): preg = vf.create_mask(pattern=pl.MaskPattern.ALL, dtype=pl.DT_UINT32) reg_src = vf.load_align(src_tile, 0) reg_shift = vf.load_align(shift_tile, 0) reg_out = vf.shift_left(reg_src, reg_shift, preg) vf.store_align(dst_tile, reg_out, preg)

注意此时 Kernel 中需要为 src 与 shift 分别声明不同数据类型的 TileType:示例里tf_u32 = pl.TileType(shape=[1, 64], dtype=pl.DT_UINT32, ...)对应被移位数,tf_i32 = pl.TileType(shape=[1, 64], dtype=pl.DT_INT32, ...)对应移位量,二者地址分别为0x00x100。Host 端用shift = torch.full([1, 64], 4, ...)构造全 4 的移位量,期望out == a << 4(右移示例为a >> 4)。

INT64 移位示例vf.shift_left(reg_a, 2, preg)/vf.shift_right(reg_a, 2, preg),tile 形状[1, 32],Host 端分别校验out == a << 2out == a >> 2

五、常见问题与编码建议

  1. 平台适配:6 个接口均仅支持 Ascend 950PR/950DT;A2/A3 平台请改用其他位运算方案(如 tile 级and_/xor,见 tile_computation/elementwise/and_.md),或先通过产品形态判断指令可用性。
  2. 数据类型一致性src0/src1/dst必须同类型;not_不支持 BF16 与 FP8 系列;移位接口仅支持 8 种整型,且 reg_tensor 模式下的 shift 必须是有符号整型。
  3. 掩码用法:所有接口都接受mask_reg操作数,且返回值也可能是mask_reg——当结果要继续作为谓词驱动其他运算时,注意其与reg_tensor的类型转换边界。
  4. 边界与未定义行为:移位量大于位宽时按文档约定输出(左移 0;右移无符号 0、有符号随符号位为 -1/0);负数移位量行为未定义,必须由调用方规避。
  5. 测试对照:文档示例均给出 Host 端 PyTorch 期望值(torch.bitwise_xor~aa << 4等),可直接作为算子正确性的 Golden 校验标准;INT64 示例同时验证了宽位宽下寄存器资源布局(tile 元素数减半)的正确性。

六、延伸阅读

  • 数据类型与寄存器结构:reg_tensor、mask_reg、DataType
  • 掩码相关操作:create_mask、mask_gen_with_reg_tensor
  • 合并模式枚举:MergeMode
  • 比较与选择指令(掩码的主要来源):comparison_and_selection/index.md
  • 底层声明源码:python/pypto_pro/language/_vf_api.py(and_or_xorshift_leftshift_rightnot_均声明于此)

以上接口文档原文位于 logical_computation 目录,包含and_not_or_shift_leftshift_rightxor六个子页面,每个页面均提供完整可复制的调用示例与测试用例。

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

【免费下载链接】pypto

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

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

相关推荐

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

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

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

立即咨询