PyPTO 掩码操作(Mask Operations)实战指南:create_mask / mask_gen_with_reg_tensor / update_mask 深入解析
人工智能编译器模型编译高性能计算深度学习CANN【免费下载链接】pyptoPyPTO发音: pai p-t-oParallel Tensor/Tile Operation编程范式。项目地址https://gitcode.com/cann/pypto点击查看免费下载导读本文围绕 CANN/PyPTO 的 SIMD 向量编程范式pypto_pro.language中三类核心掩码操作接口——vf.create_mask、vf.mask_gen_with_reg_tensor与vf.update_mask——展开完整讲解。掩码寄存器mask_reg是 VF 运算中控制元素级有效性的专用寄存器直接决定每个数据元素是否参与向量运算是处理尾块tail、条件选择select、交替筛选等高频场景的基础设施。读完本文你将掌握mask_reg的位宽与粒度机制、三类掩码生成接口的用法与约束并能基于仓库自带的调用示例与测试用例在自己的 Tile 内核中正确生成和使用掩码。一、背景PyPTO SIMD API 与掩码操作的位置PyPTOParallel 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_align、vf.store_align、vf.select、vf.abs等运算共同组成 VF 指令体系。产品支持情况三个接口目前仅支持 Ascend 950PR / Ascend 950DTAtlas A3 训练/推理系列、Atlas A2 训练/推理系列均不支持。编写可移植内核时需先确认目标硬件。二、mask_reg 工作原理256 bit 固定位宽与 dtype 粒度要正确使用三类掩码接口首先必须理解mask_reg的底层语义。VF 算子如vf.add、vf.mul执行时会根据mask_reg中每个元素对应的比特位决定该元素是否参与运算比特位为 1有效该元素参与运算结果写入目的寄存器对应位置比特位为 0无效该元素不参与运算目的寄存器对应位置置零vf.add、vf.max、vf.min、vf.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 bitb8 粒度256 bitDT_FP16 / DT_UINT16 / DT_BF1616 bit1282 bitb16 粒度256 bitDT_FP32 / DT_INT32 / DT_UINT3232 bit644 bitb32 粒度256 bitDT_INT64 / DT_UINT6464 bit328 bitb64 粒度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 的典型使用场景全量运算patternALL所有元素参与运算最常用。尾块处理当数据长度不是寄存器宽度的整数倍时用VL1~VL128限制最后一块的参与元素数。条件选择通过vf.eq、vf.gt等比较算子生成掩码再用vf.select按掩码选择元素。交替处理用H、Q、M3、M4等模式对寄存器中的部分元素进行筛选运算。以 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) - preg3.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(dtypepl.DT_FP16)pattern 默认 ALLpreg vf.create_mask(patternpl.MaskPattern.VL8)dtype 默认 FP32preg vf.create_mask()两者均取默认值。从源码注释可以推断INT64/UINT64被当作 b64 掩码宽度处理内部使用pset_b32 punpack实现每元素 2 bit 的粒度匹配所有 b8/b4 类型含 FP8E4M3FN/FP8E5M2/FP8E8M0/HF8/FP4E2M1/FP4E1M2均按 b8 掩码宽度处理。3.3 约束说明dtype 与 MaskPattern 完整对照表 1dtype 对应数据类型掩码说明dtype元素位宽元素个数每元素掩码位数总掩码位数DT_INT8 / DT_UINT8 / DT_FP8E4M3FN / DT_FP8E5M2 / DT_FP8E8M0 / DT_HF8 / DT_FP4E2M1 / DT_FP4E1M28 bit2561 bitb8 粒度256 bitDT_FP16 / DT_UINT16 / DT_BF1616 bit1282 bitb16 粒度256 bitDT_FP32 / DT_INT32 / DT_UINT3232 bit644 bitb32 粒度256 bitDT_INT64 / DT_UINT6464 bit328 bitb64 粒度256 bit表 2MaskPattern 模式说明示意以 DT_FP32 / 64 元素为例取值含义示意pypto_pro.language.MaskPattern.ALL所有元素有效1111111111111111...1111全 1pypto_pro.language.MaskPattern.ALLF所有元素无效0000000000000000...0000全 0pypto_pro.language.MaskPattern.VL1最低 1 个元素有效1000000000000000...0000pypto_pro.language.MaskPattern.VL2最低 2 个元素有效1100000000000000...0000pypto_pro.language.MaskPattern.VL4最低 4 个元素有效1111000000000000...0000pypto_pro.language.MaskPattern.VL8最低 8 个元素有效1111111100000000...0000pypto_pro.language.MaskPattern.VL16最低 16 个元素有效前 16 个 1其余 0pypto_pro.language.MaskPattern.VL32最低 32 个元素有效前 32 个 1其余 0pypto_pro.language.MaskPattern.VL64最低 64 个元素有效前 64 个 1其余 0pypto_pro.language.MaskPattern.VL128最低 128 个元素有效全部有效仅 8 位宽/16 位宽粒度下有意义pypto_pro.language.MaskPattern.H最低一半元素有效前 32 个 1后 32 个 064 元素时pypto_pro.language.MaskPattern.Q最低四分之一元素有效前 16 个 1后 48 个 064 元素时pypto_pro.language.MaskPattern.M33 的倍数位置有效每第 3 个元素为 1pypto_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 掩码并完成一次加载 → 存储的数据搬运可复制运行需要torch、torch_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(patternpl.MaskPattern.ALL, dtypepl.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], dtypepl.DT_FP32, target_memorypl.MemorySpace.Vec) in_a_grp pl.make_tile_group(typetf, addrs0x0, mutex_ids[0]) in_a in_a_grp.current() t_out_grp pl.make_tile_group(typetf, addrs0x100, 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 fnpu:{device_id} core_nums 1 torch.npu.set_device(device) a torch.randn([1, 64], devicedevice, dtypetorch.float32) out torch.empty([1, 64], devicedevice, dtypetorch.float32) example_kernelNone, core_nums torch.npu.synchronize() torch.testing.assert_close(out, a, rtol1e-5, atol1e-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_regmask_gen_with_reg_tensor(src, offset: Optional[int] None) - dst从源码注释python/pypto_pro/language/_vf_api.py可以确认该接口底层对应movvp指令将寄存器元素中的某个 bit 转换为掩码谓词。其语义为reg_tensor256B被划分为若干个 DataBlockoffset参数指定从哪个 DataBlock 生成 mask_reg。每个 DataBlock 中的每个 bit 会被 broadcast 到 mask_reg 中对应的多个 bit 位broadcast 倍数由数据类型位宽决定b16 数据类型DT_FP16、DT_BF16、DT_INT16、DT_UINT16RegTensor 划分为 16 个 DataBlock每个 16B每个 bit broadcast 到 2 bit生成 32B 的 mask_reg。offset 取值范围为 [0, 15]。b32 数据类型DT_FP32、DT_INT32、DT_UINT32RegTensor 划分为 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 DataBlock32 位宽数据类型时取值范围为 [0, 31]reg_tensor 256B 划分为 32 个 8B DataBlock。4.3 返回值返回 dst 目的操作数mask_reg。生成的 mask_reg 仅最低位有效16 位宽数据类型时每 2 bit 中仅最低位有效32 位宽数据类型时每 4 bit 中仅最低位有效。4.4 调用示例与真实测试用法以下示例演示从寄存器张量生成掩码后用于 store 控制源文档示例dtypeDT_UINT32对应 b32 粒度64 元素 × 4 bitimport 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, offset0) 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], dtypepl.DT_UINT32, target_memorypl.MemorySpace.Vec) in_a_grp pl.make_tile_group(typetf, addrs0x0, mutex_ids[0]) in_a in_a_grp.current() t_out_grp pl.make_tile_group(typetf, addrs0x100, 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 fnpu:{device_id} core_nums 1 torch.npu.set_device(device) a torch.full([1, 64], -1, devicedevice, dtypetorch.int32) out torch.empty([1, 64], devicedevice, dtypetorch.int32) example_kernelNone, core_nums torch.npu.synchronize() torch.testing.assert_close(out, a, rtol0, atol0) if __name__ __main__: test_example() print(PASSED)仓库测试 test_vf_basic_ops.py 给出了该接口更贴近实战的组合用法先用vf.create_mask(patternALL)建立全量掩码再用vf.mask_gen_with_reg_tensor(reg_u32, offset0)从 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 bitb8 粒度256 bitDT_FP16 / DT_UINT16 / DT_BF1616 bit1282 bitb16 粒度256 bitDT_FP32 / DT_INT32 / DT_UINT3232 bit644 bitb32 粒度256 bitDT_INT64 / DT_UINT6464 bit328 bitb64 粒度256 bit注意FP8 类型FP8E4M3FN/FP8E5M2/FP8E8M0/HF8和 FP4 类型FP4E2M1/FP4E1M2均为 b8 存储按 b8 粒度处理。掩码寄存器始终为 mask_reg 类型。5.4 返回值返回 preg 目标 mask_reg。5.5 调用示例与真实测试用法以下示例演示用0xFFFFFFFF32 个比特位全 1在 FP16b16 粒度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, dtypepl.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], dtypepl.DT_FP16, target_memorypl.MemorySpace.Vec) in_a_grp pl.make_tile_group(typetf, addrs0x0, mutex_ids[0]) in_a in_a_grp.current() t_out_grp pl.make_tile_group(typetf, addrs0x100, 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 fnpu:{device_id} core_nums 1 torch.npu.set_device(device) a torch.randn([1, 128], devicedevice, dtypetorch.float16) out torch.empty([1, 128], devicedevice, dtypetorch.float16) example_kernelNone, core_nums torch.npu.synchronize() torch.testing.assert_close(out, a, rtol1e-5, atol1e-5) if __name__ __main__: test_example() print(PASSED)仓库测试 test_vf_basic_ops.py 展示了update_mask的典型尾块/部分处理用法先用vf.create_mask(patternALL)建立全量掩码再vf.update_mask(8, dtypepl.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 的比特位将数据驱动的条件如符号位、比较结果转为掩码offsetb16 为 [0,15]、b32 为 [0,31]支持 6 种 16/32 位宽类型movvpvf.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负责运行时收缩。七、注意事项与易错点粒度不等于类型dtype只决定掩码粒度每元素占多少 bitmask_reg 类型始终不变。混淆这一点是新手最容易出错的地方。VL128 的适用前提VL128仅在 8 位宽/16 位宽粒度下有意义b32/b64 粒度下元素数不足 128无法使用。offset 越界mask_gen_with_reg_tensor的 offset 范围与位宽强相关——b16 最大 15、b32 最大 31越界即非法。FP8/FP4 的粒度归并FP8E4M3FN、FP8E5M2、FP8E8M0、HF8 以及 FP4E2M1、FP4E1M2 均按 b8 存储与粒度处理。mask_reg 数量上限MaskReg 寄存器上限为 16编译器自动复用生命周期结束的寄存器长时间持有大量掩码可能触发寄存器压力。硬件支持范围三类接口目前仅支持 Ascend 950PR / Ascend 950DT编写跨平台内核前需先校验目标产品。掩码位为 0 的语义无效元素的目的是寄存器对应位置置零少数算子可通过mode保留原值这是设计掩码运算逻辑时必须牢记的语义差异。八、参考资源接口文档create_mask、mask_gen_with_reg_tensor、update_mask类型与配套mask_reg.md、MaskPattern.md源码声明python/pypto_pro/language/_vf_api.pycreate_mask/update_mask、python/pypto_pro/language/_vf_api.pymask_gen_with_reg_tensor参数解析python/pypto_pro/language/parser/_call_parser.pypattern 等枚举关键字约束测试用例python/tests/st/pypto_pro/frontend/vf_api/test_vf_basic_ops.py含 mask_gen 组合用法、update_mask 尾块用法等赞分享人工智能编译器模型编译高性能计算深度学习CANN【免费下载链接】pyptoPyPTO发音: pai p-t-oParallel Tensor/Tile Operation编程范式。项目地址https://gitcode.com/cann/pypto点击查看免费下载相关推荐深入解析Text Mask跨框架输入掩码解决方案深入解析Text Mask跨框架输入掩码解决方案 Text Mask是一个功能强大的输入掩码库专门设计用于处理表单输入字段的格式化需求。该项目采用模块化架构前端Mask R-CNN 深度掩码头DeepMAC实战tensorflow/models 中的深度掩码分割模型完全解析Mask R CNN 深度掩码头DeepMAC实战tensorflow/models 中的深度掩码分割模型完全解析 Mask R CNN 的实例掩码质量长人工智能深度学习计算机视觉NLP语音深入解析 PyPTO 的 pypto.clip 数据裁剪操作用法、约束与源码实现深入解析 PyPTO 的 pypto.clip 数据裁剪操作用法、约束与源码实现 本文以 PyPTO 的张量算子 API pypto.clip 为主线完整讲人工智能编译器模型编译高性能计算深度学习CANN创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考