PyPTO vf.create_mask 详解:VF 掩码寄存器创建与元素级运算控制
2026/9/20 21:37:46 网站建设 项目流程
  • 人工智能
  • 编译器
  • 模型编译
  • 高性能计算
  • 深度学习
  • CANN

【免费下载链接】pypto

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

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

导读

vf.create_mask是 PyPTO(Parallel Tensor/Tile Operation 编程范式)中 Vector Function(VF)指令域的掩码寄存器创建接口,用于生成控制后续 VF 运算(如vf.addvf.mulvf.store_align等)元素级有效性的 mask_reg。本文基于 create_mask 接口文档 并结合仓库源码,系统讲解 mask_reg 的位宽粒度原理、MaskPattern 各模式语义、参数约束、在 astype 精度转换等场景中的行为,以及可运行的完整调用示例,帮助你准确掌握在 PyPTO 内核中按需筛选 VF 运算元素的方法。

mask_reg 的工作原理

mask_reg 是 VF 运算中控制元素级有效性的专用寄存器。VF 算子(如vf.addvf.mul等)在执行时,会根据 mask_reg 中每个元素对应的比特位决定该元素是否参与运算:

  • 比特位为 1(有效):该元素参与运算,结果写入目的寄存器对应位置。
  • 比特位为 0(无效):该元素不参与运算,目的寄存器对应位置置零。vf.addvf.maxvf.minvf.full等少数算子支持通过mode参数选择保留原值。

mask_reg 的总位宽固定为256 bit,其粒度由dtype参数决定;每个数据元素对应的掩码位数随元素位宽变化。例如:

dtype元素位宽元素个数每元素掩码位数总掩码位数
DT_INT8 / DT_UINT88 bit2561 bit(8 位宽粒度)256 bit
DT_FP16 / DT_UINT16 / DT_BF1616 bit1282 bit(16 位宽粒度)256 bit
DT_FP32 / DT_INT32 / DT_UINT3232 bit644 bit(32 位宽粒度)256 bit
DT_INT64 / DT_UINT6464 bit328 bit(64 位宽粒度)256 bit

[!CAUTION] 注意dtype参数决定的是掩码粒度(即 mask_reg 中每多少个 bit 对应一个数据元素),而非 mask_reg 本身的类型。mask_reg 类型始终不变。

mask_reg 的典型使用场景

  1. 全量运算pattern=ALL,所有元素参与运算(最常用)。
  2. 尾块处理:当数据长度不是寄存器宽度的整数倍时,用 VL1~VL128 限制最后一块的参与元素数。
  3. 条件选择:通过vf.eqvf.gt等比较算子生成掩码,再用vf.select按掩码选择元素。
  4. 交替处理:用 H、Q、M3、M4 等模式对寄存器中的部分元素进行筛选运算。

以 b8 数据类型为例,不同 MaskPattern 模式下 CreateMask 接口的元素选取如下图所示:

图 1b8 数据类型下 CreateMask 接口不同 MaskPattern 模式下元素选取

astype 精度转换中的 mask_reg

不同数据类型下元素对应的 mask 位宽不一致,在 astype 进行类型转换时,mask_reg 根据输入的源操作数进行有效元素筛选。mask_reg 和 RegLayout 同时作用时,16 位宽与 32 位宽的相互转换过程如下图所示:

图 2astype 16 位宽到 32 位宽类型转换过程

图 3astype 32 位宽到 16 位宽类型转换过程

函数原型

create_mask(pattern: Optional[MaskPattern] = None, dtype: Optional[DType] = None) -> preg

参数说明

参数输入/输出说明
pattern输入可选,掩码模式,pattern 参数决定 mask_reg 中哪些元素被设置为有效(1),哪些被设置为无效(0),对应 MaskPattern 类型。支持的模式见约束说明,默认pypto_pro.language.MaskPattern.ALL
dtype输入可选,掩码对应的数据类型,决定掩码粒度(即每多少 bit 对应一个数据元素)。如pypto_pro.language.DT_FP32对应 32 位宽粒度(64 元素 × 4 bit),全部对应关系请见约束说明。掩码寄存器总位宽固定为 256 bit,默认pypto_pro.language.DT_FP32

约束说明

dtype 与掩码粒度对应关系

表 1dtype 对应数据类型掩码说明

dtype元素位宽元素个数每元素掩码位数总掩码位数
DT_INT8 / DT_UINT8 / DT_FP8E4M3FN / DT_FP8E5M2 / DT_FP8E8M0 / DT_HF8 / DT_FP4E2M1 / DT_FP4E1M28 bit2561 bit(b8 粒度)256 bit
DT_FP16 / DT_UINT16 / DT_BF1616 bit1282 bit(b16 粒度)256 bit
DT_FP32 / DT_INT32 / DT_UINT3232 bit644 bit(b32 粒度)256 bit
DT_INT64 / DT_UINT6464 bit328 bit(b64 粒度)256 bit

仓库实现说明:在 VF API 声明 中,create_mask的 docstring 明确指出,所有 b8/b4 类型(如 FP8E4M3FN、FP8E5M2、FP8E8M0、HF8、FP4E2M1、FP4E1M2 等)统一按 b8 掩码宽度处理;INT64/UINT64 按 b64 掩码宽度处理,内部通过pset_b32+punpack实现每元素 2 bit 的粒度匹配。

MaskPattern 模式说明

表 2MaskPattern 模式说明

取值含义示意(以 DT_FP32 / 64 元素为例)
pypto_pro.language.MaskPattern.ALL所有元素有效1111111111111111...1111(全 1)
pypto_pro.language.MaskPattern.ALLF所有元素无效0000000000000000...0000(全 0)
pypto_pro.language.MaskPattern.VL1最低 1 个元素有效1000000000000000...0000
pypto_pro.language.MaskPattern.VL2最低 2 个元素有效1100000000000000...0000
pypto_pro.language.MaskPattern.VL4最低 4 个元素有效1111000000000000...0000
pypto_pro.language.MaskPattern.VL8最低 8 个元素有效1111111100000000...0000
pypto_pro.language.MaskPattern.VL16最低 16 个元素有效前 16 个 1,其余 0
pypto_pro.language.MaskPattern.VL32最低 32 个元素有效前 32 个 1,其余 0
pypto_pro.language.MaskPattern.VL64最低 64 个元素有效前 64 个 1,其余 0
pypto_pro.language.MaskPattern.VL128最低 128 个元素有效全部有效(仅 8 位宽/16 位宽粒度下有意义)
pypto_pro.language.MaskPattern.H最低一半元素有效前 32 个 1,后 32 个 0(64 元素时)
pypto_pro.language.MaskPattern.Q最低四分之一元素有效前 16 个 1,后 48 个 0(64 元素时)
pypto_pro.language.MaskPattern.M33 的倍数位置有效每第 3 个元素为 1
pypto_pro.language.MaskPattern.M44 的倍数位置有效每第 4 个元素为 1

返回值说明

返回 preg 目标 mask_reg。

调用示例

下面是一个完整的可运行示例,演示create_maskload_align/store_align配合实现全量元素搬运:

import os import pypto_pro.language as pl import torch import torch_npu @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) vf.store_align(dst_tile, reg, preg) @pl.jit() def example_kernel( a: 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() t_out_grp = pl.make_tile_group(type=tf, addrs=0x100, mutex_ids=[1]) t_out = t_out_grp.current() with pl.section_vector(): pl.load(in_a, a, [0, 0]) example_vf(in_a, 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) out = torch.empty([1, 64], device=device, dtype=torch.float32) example_kernelNone, core_nums torch.npu.synchronize() torch.testing.assert_close(out, a, rtol=1e-5, atol=1e-5) if __name__ == "__main__": test_example() print("PASSED")

示例要点解析

  • VF 函数与内核分离example_vf使用@pl.vector_function装饰,内部执行 VF 指令;example_kernel使用@pl.jit()装饰,负责 tile 的分配、加载与存储。
  • 掩码贯穿加载与存储vf.load_alignvf.store_align都接收preg作为掩码参数,pattern=ALL 表示全量元素有效,此时运算等价于直接搬运。
  • 运行环境:示例依赖torch_npu及 NPU 设备,通过TILE_FWK_DEVICE_ID环境变量指定设备号(默认 0)。
  • 结果校验:使用torch.testing.assert_close对比输出与输入张量,误差阈值 rtol=1e-5、atol=1e-5。

源码级实现佐证

双默认值设计

从 VF API 声明 可以看出,create_mask的两个参数均可独立缺省,并且内部实现支持以下三种调用形态:

preg = vf.create_mask(dtype=pl.DT_FP16) # pattern 默认 ALL preg = vf.create_mask(pattern=pl.MaskPattern.VL8) # dtype 默认 FP32 preg = vf.create_mask() # 两者均取默认值

MaskReg 寄存器类型识别

在 调用解析器 中,create_mask被登记为_VF_MASK_PRODUCING_OPS集合中的一员(与update_maskget_mask_sprmask_gen_with_reg_tensor并列),即其返回值为 MaskReg 而非普通 RegTensor。该信息被赋值解析器用于跟踪 MaskReg 变量,进而影响selectmoveand_or_等统一操作(_VF_UNIFIED_OPS)的目的寄存器类型推断:当源操作数中出现已知的 MaskReg 变量时,目的寄存器会被声明为 MaskReg。

MaskPattern 枚举定义

MaskPattern 枚举在 IR 绑定层 中注册到 Python 侧,除文档列出的 ALL、ALLF、VL1~VL128、M3、M4、H、Q 外,还额外包含VL3模式(最低 3 个元素有效),共 16 个枚举值,可直接通过pypto_pro.language.MaskPattern访问(见 language 包导出)。

产品支持情况

产品型号支持情况
Ascend 950PR / Ascend 950DT支持
Atlas A3 训练系列产品 / Atlas A3 推理系列产品不支持
Atlas A2 训练系列产品 / Atlas A2 推理系列产品不支持

说明:create_mask属于 A5 架构的 VF(Vector Function)单元指令,VF API 声明 中明确标注其仅在@pl.vector_function装饰的函数内使用,Compute 类算子必须使用赋值形式dst = vf.xxx(...),只有 store 等副作用算子才以裸语句形式调用。

相关阅读

  • mask_reg 接口文档
  • MaskPattern 类型文档
  • VF API 声明源码
  • VF 调用解析器
  • VF 赋值解析器
  • SIMD-API 索引
  • 人工智能
  • 编译器
  • 模型编译
  • 高性能计算
  • 深度学习
  • CANN

【免费下载链接】pypto

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

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

相关推荐

上一篇:揭秘terraform-provider-libvirt数据来源:节点信息与设备管理技巧
下一篇:PySlowFast错误调试手册:常见训练问题解决方案

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

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

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

立即咨询