PyPTO 掩码操作(Mask Operations)实战指南:create_mask / mask_gen_with_reg_tensor / update_mask 深入解析
2026/9/20 4:31:57 网站建设 项目流程
  • 人工智能
  • 编译器
  • 模型编译
  • 高性能计算
  • 深度学习
  • CANN

【免费下载链接】pypto

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

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

导读

本文围绕 CANN/PyPTO 的 SIMD 向量编程范式(pypto_pro.language)中三类核心掩码操作接口——vf.create_maskvf.mask_gen_with_reg_tensorvf.update_mask——展开完整讲解。掩码寄存器(mask_reg)是 VF 运算中控制元素级有效性的专用寄存器,直接决定每个数据元素是否参与向量运算,是处理尾块(tail)、条件选择(select)、交替筛选等高频场景的基础设施。读完本文,你将掌握mask_reg的位宽与粒度机制、三类掩码生成接口的用法与约束,并能基于仓库自带的调用示例与测试用例,在自己的 Tile 内核中正确生成和使用掩码。

一、背景:PyPTO SIMD API 与掩码操作的位置

PyPTO(Parallel Tensor/Tile Operation 编程范式)的 SIMD-API 按功能划分为基础数据结构、缓存控制、控制流、Cube 计算、数据搬运、量化、寄存器计算(reg_computation)、资源管理、同步、系统变量、Tile 计算、转置与元素访问、工具等多个目录。其中docs/zh/api/pro_api/SIMD-API/reg_computation/mask_operations/目录专门收拢掩码相关操作,包含三个接口文档:

  • create_mask:按固定模式创建掩码;
  • mask_gen_with_reg_tensor:从寄存器张量的比特位生成掩码;
  • update_mask:从标量值更新掩码。

三者均作用于mask_reg寄存器。关于mask_reg本身的完整定义(原型、参数、约束)可参考 mask_reg.md,其配套的类型枚举 MaskPattern.md 则定义了create_mask支持的掩码模式。此外,在 Python 源码 的vfAPI 声明中可以看到三者与vf.load_alignvf.store_alignvf.selectvf.abs等运算共同组成 VF 指令体系。

产品支持情况:三个接口目前仅支持 Ascend 950PR / Ascend 950DT;Atlas A3 训练/推理系列、Atlas A2 训练/推理系列均不支持。编写可移植内核时需先确认目标硬件。

二、mask_reg 工作原理:256 bit 固定位宽与 dtype 粒度

要正确使用三类掩码接口,首先必须理解mask_reg的底层语义。VF 算子(如vf.addvf.mul)执行时会根据mask_reg中每个元素对应的比特位决定该元素是否参与运算:

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

mask_reg的总位宽固定为256 bit,但其粒度由关联的dtype参数决定:每个数据元素对应的掩码位数随元素位宽变化。不同 dtype 的对应关系如下:

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

[!CAUTION] 注意dtype参数决定的是掩码粒度(即 mask_reg 中每多少个 bit 对应一个数据元素),而非 mask_reg 本身的类型。mask_reg 类型始终不变。FP8 类型(FP8E4M3FN/FP8E5M2/FP8E8M0/HF8)和 FP4 类型(FP4E2M1/FP4E1M2)均为 b8 存储,按 b8 粒度处理。这一点在 create_mask、update_mask 和 mask_reg 三处文档中均有明确说明。

mask_reg 的典型使用场景

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

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

astype 精度转换中的 mask_reg

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

三、vf.create_mask:按固定模式创建掩码

3.1 功能与函数原型

vf.create_mask用于创建 mask_reg,指定参与后续 VF 运算的元素范围:

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

3.2 参数说明

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

在源码中,create_mask声明于 python/pypto_pro/language/_vf_api.py,两个 kwargs 均可选且可独立指定,例如:

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

从源码注释可以推断:INT64/UINT64被当作 b64 掩码宽度处理,内部使用pset_b32 + punpack实现每元素 2 bit 的粒度匹配;所有 b8/b4 类型(含 FP8E4M3FN/FP8E5M2/FP8E8M0/HF8/FP4E2M1/FP4E1M2)均按 b8 掩码宽度处理。

3.3 约束说明:dtype 与 MaskPattern 完整对照

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

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

表 2:MaskPattern 模式说明(示意以 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

完整的枚举定义见 MaskPattern.md:除上述取值外还包含VL3(最低 3 个元素有效),其语义为“每 3 个元素中第 1 个有效”(M3)、"每 4 个元素中第 1 个有效"(M4)、"低半部分有效"(H)、"低四分之一有效"(Q)。在 Python 前端中,MaskPattern由 python/pypto_pro/language/init.py 从pypto.ir导出,并由 call_parser.py 在解析 VF 调用时将pattern关键字参数约束为MaskPattern枚举类型。

3.4 返回值

返回 preg 目标 mask_reg。vf.mask_reg本身不能直接调用,由编译器在赋值形式中自动声明(如preg = vf.create_mask(...)),且在@pl.vector_function函数内创建和使用、函数结束后自动释放;MaskReg 寄存器数量上限为 16,编译器会自动复用生命周期结束的寄存器与预留内存,若两者均存在可用空间则优先复用寄存器。

3.5 调用示例

以下完整示例演示了在 Tile 内核中创建 ALL 掩码并完成一次"加载 → 存储"的数据搬运,可复制运行(需要torchtorch_npu以及支持 950 系列硬件的运行环境):

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")

代码要点:pl.section_vector()划定向量执行段;pl.load/pl.store负责 HBM 与 Vec 内存之间的 Tile 数据搬运;vf.load_align将对齐地址的 Tile 数据加载为寄存器张量,vf.store_align在掩码preg控制下写回目的 Tile。同样的 "ALL 掩码 + 搬运" 模式在仓库测试 test_vf_basic_ops.py 中大量出现(如_vf_kernel_49_truncate_maskgen_0_vf_kernel_65_update_mask_0等),可作为回归验证参考。

四、vf.mask_gen_with_reg_tensor:从寄存器比特生成掩码

4.1 功能与函数原型

vf.mask_gen_with_reg_tensor从 reg_tensor 的指定数据块(DataBlock)的 bit 位生成 mask_reg:

mask_gen_with_reg_tensor(src, offset: Optional[int] = None) -> dst

从源码注释(python/pypto_pro/language/_vf_api.py)可以确认,该接口底层对应movvp指令:将寄存器元素中的某个 bit 转换为掩码谓词。其语义为:reg_tensor(256B)被划分为若干个 DataBlock,offset参数指定从哪个 DataBlock 生成 mask_reg。每个 DataBlock 中的每个 bit 会被 broadcast 到 mask_reg 中对应的多个 bit 位,broadcast 倍数由数据类型位宽决定:

  • b16 数据类型(DT_FP16、DT_BF16、DT_INT16、DT_UINT16):RegTensor 划分为 16 个 DataBlock(每个 16B),每个 bit broadcast 到 2 bit,生成 32B 的 mask_reg。offset 取值范围为 [0, 15]。
  • b32 数据类型(DT_FP32、DT_INT32、DT_UINT32):RegTensor 划分为 32 个 DataBlock(每个 8B),每个 bit broadcast 到 4 bit,生成 32B 的 mask_reg。offset 取值范围为 [0, 31]。

b16 与 b32 两种数据类型下的搬运原理分别如下图所示:

4.2 参数与约束

参数输入/输出说明
src输入源操作数,reg_tensor。支持的数据类型为:DT_FP16、DT_BF16、DT_INT16、DT_UINT16、DT_FP32、DT_INT32、DT_UINT32。
offset输入可选,指定从 src 的哪个 DataBlock 生成 mask_reg,默认 0。16 位宽数据类型时取值范围为 [0, 15](reg_tensor 256B 划分为 16 个 16B DataBlock);32 位宽数据类型时取值范围为 [0, 31](reg_tensor 256B 划分为 32 个 8B DataBlock)。

4.3 返回值

返回 dst 目的操作数,mask_reg。生成的 mask_reg 仅最低位有效:16 位宽数据类型时每 2 bit 中仅最低位有效,32 位宽数据类型时每 4 bit 中仅最低位有效。

4.4 调用示例与真实测试用法

以下示例演示从寄存器张量生成掩码后用于 store 控制(源文档示例,dtype=DT_UINT32对应 b32 粒度,64 元素 × 4 bit):

import os import pypto_pro.language as pl import torch import torch_npu @pl.vector_function def example_vf(src_tile, dst_tile): reg = vf.load_align(src_tile, 0) dst = vf.mask_gen_with_reg_tensor(reg, offset=0) vf.store_align(dst_tile, reg, dst) @pl.jit() def example_kernel( a: pl.Tensor[[pl.DYNAMIC, pl.DYNAMIC], pl.DT_UINT32], out: pl.Tensor[[pl.DYNAMIC, pl.DYNAMIC], pl.DT_UINT32], ): tf = pl.TileType(shape=[1, 64], dtype=pl.DT_UINT32, 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.full([1, 64], -1, device=device, dtype=torch.int32) out = torch.empty([1, 64], device=device, dtype=torch.int32) example_kernelNone, core_nums torch.npu.synchronize() torch.testing.assert_close(out, a, rtol=0, atol=0) if __name__ == "__main__": test_example() print("PASSED")

仓库测试 test_vf_basic_ops.py 给出了该接口更贴近实战的组合用法:先用vf.create_mask(pattern=ALL)建立全量掩码,再用vf.mask_gen_with_reg_tensor(reg_u32, offset=0)从 UINT32 寄存器的 bit 0 生成条件掩码,随后通过vf.select(reg_a, reg_b, gen_mask)完成按掩码的元素级选择,最后在 ALL 掩码下vf.store_align写回——即"掩码生成 → 掩码驱动运算"的完整链路。

五、vf.update_mask:从标量值更新掩码

5.1 功能与函数原型

vf.update_mask从标量值更新 mask_reg,根据当前scalarValue的值生成对应长度的有效位掩码:

update_mask(scalar, dtype: Optional[DType] = None) -> preg

以 16 位宽数据类型为例,掩码生成过程如下图所示:

5.2 参数说明

参数输入/输出说明
scalar输入标量值,其比特位定义新的掩码模式。
dtype输入可选,掩码对应的数据类型,决定掩码宽度(默认pypto_pro.language.DT_FP32)。本接口操作数为寄存器,不涉及地址对齐;本接口不修改全局寄存器的值。

从源码声明(python/pypto_pro/language/_vf_api.py)可以看出:该接口将标量值的比特位直接写入掩码寄存器,dtype仅用于选择掩码宽度(默认 FP32 对应 b32)。

5.3 约束说明

dtype 参数决定掩码粒度(即每多少 bit 对应一个数据元素),掩码寄存器总位宽固定为 256 bit,对应关系与create_mask完全一致:

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

注意:FP8 类型(FP8E4M3FN/FP8E5M2/FP8E8M0/HF8)和 FP4 类型(FP4E2M1/FP4E1M2)均为 b8 存储,按 b8 粒度处理。掩码寄存器始终为 mask_reg 类型。

5.4 返回值

返回 preg 目标 mask_reg。

5.5 调用示例与真实测试用法

以下示例演示用0xFFFFFFFF(32 个比特位全 1)在 FP16(b16 粒度,128 元素 × 2 bit)下构造全有效掩码(源文档示例):

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.update_mask(0xFFFFFFFF, dtype=pl.DT_FP16) 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_FP16], out: pl.Tensor[[pl.DYNAMIC, pl.DYNAMIC], pl.DT_FP16], ): tf = pl.TileType(shape=[1, 128], dtype=pl.DT_FP16, 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, 128], device=device, dtype=torch.float16) out = torch.empty([1, 128], device=device, dtype=torch.float16) 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")

仓库测试 test_vf_basic_ops.py 展示了update_mask的典型"尾块/部分处理"用法:先用vf.create_mask(pattern=ALL)建立全量掩码,再vf.update_mask(8, dtype=pl.DT_FP32)生成只允许最低 8 个元素参与的掩码,随后在preg_tail控制下执行vf.abs取绝对值运算,最后在 ALL 掩码下写回——这正是"按需收缩有效元素范围"的标准套路。

六、三类接口的对比与选型建议

接口掩码来源典型用途关键参数底层指令(源码注释)
vf.create_mask固定模式枚举(ALL/ALLF/VL*/H/Q/M3/M4)全量运算、尾块、交替筛选pattern(默认 ALL)、dtype(默认 FP32)按模式初始化谓词寄存器
vf.mask_gen_with_reg_tensor寄存器张量 DataBlock 的比特位将数据驱动的条件(如符号位、比较结果)转为掩码offset(b16 为 [0,15]、b32 为 [0,31]),支持 6 种 16/32 位宽类型movvp
vf.update_mask标量值的比特位编译期/运行时已知的固定有效区间、尾块收缩scalar(比特位定义掩码)、dtype(默认 FP32)按标量写掩码寄存器

选型建议:

  • 需要"所有元素参与"或"固定比例参与"时,优先vf.create_mask(默认值即 ALL,最省心);
  • 数据长度不固定、需要在运行时决定参与元素数时,用vf.update_mask(scalar, dtype=...)动态收缩;
  • 掩码本身由数据内容决定(例如某个 bit 是否为 1)时,用vf.mask_gen_with_reg_tensor将数据位 broadcast 成掩码,再配合vf.select实现条件选择。

三者生成的 mask_reg 粒度语义一致(dtype 决定每元素掩码位数,总位宽 256 bit),因此可以互相配合使用:create_mask负责建立基准掩码,mask_gen_with_reg_tensor负责数据驱动掩码,update_mask负责运行时收缩。

七、注意事项与易错点

  1. 粒度不等于类型dtype只决定掩码粒度(每元素占多少 bit),mask_reg 类型始终不变。混淆这一点是新手最容易出错的地方。
  2. VL128 的适用前提VL128仅在 8 位宽/16 位宽粒度下有意义;b32/b64 粒度下元素数不足 128,无法使用。
  3. offset 越界mask_gen_with_reg_tensor的 offset 范围与位宽强相关——b16 最大 15、b32 最大 31,越界即非法。
  4. FP8/FP4 的粒度归并:FP8E4M3FN、FP8E5M2、FP8E8M0、HF8 以及 FP4E2M1、FP4E1M2 均按 b8 存储与粒度处理。
  5. mask_reg 数量上限:MaskReg 寄存器上限为 16,编译器自动复用生命周期结束的寄存器;长时间持有大量掩码可能触发寄存器压力。
  6. 硬件支持范围:三类接口目前仅支持 Ascend 950PR / Ascend 950DT,编写跨平台内核前需先校验目标产品。
  7. 掩码位为 0 的语义:无效元素的目的是寄存器对应位置置零(少数算子可通过mode保留原值),这是设计掩码运算逻辑时必须牢记的语义差异。

八、参考资源

  • 接口文档:create_mask、mask_gen_with_reg_tensor、update_mask
  • 类型与配套:mask_reg.md、MaskPattern.md
  • 源码声明:python/pypto_pro/language/_vf_api.py(create_mask/update_mask)、python/pypto_pro/language/_vf_api.py(mask_gen_with_reg_tensor)
  • 参数解析:python/pypto_pro/language/parser/_call_parser.py(pattern 等枚举关键字约束)
  • 测试用例:python/tests/st/pypto_pro/frontend/vf_api/test_vf_basic_ops.py(含 mask_gen 组合用法、update_mask 尾块用法等)
  • 人工智能
  • 编译器
  • 模型编译
  • 高性能计算
  • 深度学习
  • CANN

【免费下载链接】pypto

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

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

相关推荐

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

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

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

立即咨询