- 人工智能
- 编译器
- 模型编译
- 高性能计算
- 深度学习
- CANN
【免费下载链接】pypto
PyPTO(发音: pai p-t-o):Parallel Tensor/Tile Operation编程范式。
导读
vf.create_mask是 PyPTO(Parallel Tensor/Tile Operation 编程范式)中 Vector Function(VF)指令域的掩码寄存器创建接口,用于生成控制后续 VF 运算(如vf.add、vf.mul、vf.store_align等)元素级有效性的 mask_reg。本文基于 create_mask 接口文档 并结合仓库源码,系统讲解 mask_reg 的位宽粒度原理、MaskPattern 各模式语义、参数约束、在 astype 精度转换等场景中的行为,以及可运行的完整调用示例,帮助你准确掌握在 PyPTO 内核中按需筛选 VF 运算元素的方法。
mask_reg 的工作原理
mask_reg 是 VF 运算中控制元素级有效性的专用寄存器。VF 算子(如vf.add、vf.mul等)在执行时,会根据 mask_reg 中每个元素对应的比特位决定该元素是否参与运算:
- 比特位为 1(有效):该元素参与运算,结果写入目的寄存器对应位置。
- 比特位为 0(无效):该元素不参与运算,目的寄存器对应位置置零。
vf.add、vf.max、vf.min、vf.full等少数算子支持通过mode参数选择保留原值。
mask_reg 的总位宽固定为256 bit,其粒度由dtype参数决定;每个数据元素对应的掩码位数随元素位宽变化。例如:
| dtype | 元素位宽 | 元素个数 | 每元素掩码位数 | 总掩码位数 |
|---|---|---|---|---|
| DT_INT8 / DT_UINT8 | 8 bit | 256 | 1 bit(8 位宽粒度) | 256 bit |
| DT_FP16 / DT_UINT16 / DT_BF16 | 16 bit | 128 | 2 bit(16 位宽粒度) | 256 bit |
| DT_FP32 / DT_INT32 / DT_UINT32 | 32 bit | 64 | 4 bit(32 位宽粒度) | 256 bit |
| DT_INT64 / DT_UINT64 | 64 bit | 32 | 8 bit(64 位宽粒度) | 256 bit |
[!CAUTION] 注意
dtype参数决定的是掩码粒度(即 mask_reg 中每多少个 bit 对应一个数据元素),而非 mask_reg 本身的类型。mask_reg 类型始终不变。
mask_reg 的典型使用场景
- 全量运算:
pattern=ALL,所有元素参与运算(最常用)。 - 尾块处理:当数据长度不是寄存器宽度的整数倍时,用 VL1~VL128 限制最后一块的参与元素数。
- 条件选择:通过
vf.eq、vf.gt等比较算子生成掩码,再用vf.select按掩码选择元素。 - 交替处理:用 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_FP4E1M2 | 8 bit | 256 | 1 bit(b8 粒度) | 256 bit |
| DT_FP16 / DT_UINT16 / DT_BF16 | 16 bit | 128 | 2 bit(b16 粒度) | 256 bit |
| DT_FP32 / DT_INT32 / DT_UINT32 | 32 bit | 64 | 4 bit(b32 粒度) | 256 bit |
| DT_INT64 / DT_UINT64 | 64 bit | 32 | 8 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.M3 | 3 的倍数位置有效 | 每第 3 个元素为 1 |
pypto_pro.language.MaskPattern.M4 | 4 的倍数位置有效 | 每第 4 个元素为 1 |
返回值说明
返回 preg 目标 mask_reg。
调用示例
下面是一个完整的可运行示例,演示create_mask与load_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_align与vf.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_mask、get_mask_spr、mask_gen_with_reg_tensor并列),即其返回值为 MaskReg 而非普通 RegTensor。该信息被赋值解析器用于跟踪 MaskReg 变量,进而影响select、move、and_、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编程范式。
相关推荐
PyPTO Pro vf.exp:SIMD 寄存器级指数运算接口详解(Ascend 950)
PyPTO Pro vf.exp:SIMD 寄存器级指数运算接口详解(Ascend 950) vf.exp 是 PyPTO Pro( pypto_pro )SI
人工智能编译器模型编译高性能计算深度学习CANNPyPTO Pro API:vf.log2 寄存器级以 2 为底对数运算详解与实战
PyPTO Pro API:vf.log2 寄存器级以 2 为底对数运算详解与实战 本篇聚焦 PyPTO(Parallel Tensor/Tile Operat
人工智能编译器模型编译高性能计算深度学习CANNPyPTO 逐元素最大值运算 pypto.maximum 接口详解与源码实践
PyPTO 逐元素最大值运算 pypto.maximum 接口详解与源码实践 本篇技术指南围绕 CANN PyPTO 张量运算 API 中的 pypto.max
人工智能编译器模型编译高性能计算深度学习CANN
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考