十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

CANN PyPTO 向量函数 `vf.max` 寄存器级逐元素求最大值实战指南

CANN PyPTO 向量函数 `vf.max` 寄存器级逐元素求最大值实战指南 CANN PyPTO 向量函数vf.max寄存器级逐元素求最大值实战指南【免费下载链接】pyptoPyPTO发音: pai p-t-oParallel Tensor/Tile Operation编程范式。项目地址: https://gitcode.com/cann/pypto导读vf.max是 CANN PyPTO 向量函数vector functionVF编程范式中的寄存器级按元素求最大值接口用于在向量寄存器层面完成dst_i max(src0_i, src1_i)的运算并通过掩码寄存器mask_reg精确控制参与运算的元素。本文围绕该接口系统讲解其产品支持范围、函数原型、参数语义、掩码与合并模式MergeMode的底层机制、数据类型约束以及两种可复制运行的完整调用示例FP32 与 INT64并结合仓库源码与测试用例深入说明其实现与验证方式。读完本文你将能够独立在pl.vector_function内核中编写基于寄存器与掩码的逐元素最大值运算并掌握 PyPTO VF 计算接口通用的寄存器 掩码 合并模式调用范式。接口概览寄存器级 VF 计算中的max在 PyPTO 的 VF 计算体系中运算对象不是直接的张量Tensor或 UB Tile而是向量寄存器reg_tensor。vf.max与vf.add、vf.sub、vf.mul、vf.div等接口同属二元寄存器运算一类从 UB Tile 中加载数据到寄存器在寄存器中完成逐元素运算再将结果寄存器写回 UB Tile。vf.max的功能说明与数学定义见 max.md$$dst_i \max(src0_i, src1_i)$$即根据掩码寄存器preg对源操作数src0、src1进行按元素求最大值操作将结果写入目的操作数dst。在源码层面接口声明位于 python/pypto_pro/language/_vf_api.pystaticmethod _api_decl def max(src0, src1, preg, mode: Optional[MergeMode] None): rElement-wise maximum of two source registers. For each lane i where mask[i] is active, compares the corresponding elements of src0 and src1 and writes the larger value to dst[i]. 可见其核心语义为仅对掩码有效active的 lane 执行比较并写入较大值。该接口采用赋值形式调用目的寄存器由编译器隐式声明即reg_out vf.max(reg_a, reg_b, preg)中的reg_out无需手动创建。产品支持情况与 PyPTO VF 寄存器体系reg_tensor、mask_reg、MergeMode保持一致vf.max的产品支持情况如下产品支持情况Ascend 950PR / Ascend 950DT支持Atlas A3 训练系列产品 / Atlas A3 推理系列产品不支持Atlas A2 训练系列产品 / Atlas A2 推理系列产品不支持这意味着vf.max是面向 Ascend 950 系列支持 VF 寄存器指令集的算力平台的接口在 A2/A3 系列上无法使用。这一限制同样适用于 reg_tensor、mask_reg 与 MergeMode 等配套类型。函数原型与参数说明函数原型max(src0, src1, preg, mode: Optional[MergeMode] None) - dst参数详解参数输入/输出说明src0输入源操作数 0reg_tensor源操作数 src0 与目的操作数 dst 的数据类型保持一致。支持的数据类型为DT_INT8、DT_UINT8、DT_INT16、DT_UINT16、DT_FP16、DT_BF16、DT_INT32、DT_UINT32、DT_FP32、DT_INT64、DT_UINT64。src1输入源操作数 1reg_tensor数据类型与 src0 一致。preg输入mask_reg 掩码寄存器控制哪些元素参与运算。mode输入可选对应 MergeMode 类型。pypto_pro.language.MergeMode.ZEROING默认preg 未筛选的元素在 dst 中置 0pypto_pro.language.MergeMode.MERGING 当前不支持。需要注意两点寄存器数据类型一致性src0、src1、dst三者的数据类型必须一致且掩码寄存器preg的 dtype 一般也应当与reg_tensor一致见下文掩码粒度说明。返回值是寄存器而非张量vf.max返回的是reg_tensor目的寄存器后续必须通过vf.store_align等指令写回 UB Tile才能被pl.store搬运到全局内存。源操作数的载体reg_tensor 寄存器src0、src1都是向量寄存器reg_tensor其关键特性详见 reg_tensor.md寄存器总大小固定为 256 字节不同 dtype 对应不同的元素个数dtype元素宽度元素个数DT_INT8 / DT_UINT8 / DT_HF8 / DT_FP8E4M3FN / DT_FP8E5M2 / DT_FP8E8M08 bit256DT_FP4 / DT_FP4E2M1 / DT_FP4E1M24 bitb8 打包存储2 元素/字节256逻辑 512DT_INT16 / DT_UINT16 / DT_FP16 / DT_BF1616 bit128DT_INT32 / DT_UINT32 / DT_FP3232 bit64DT_INT64 / DT_UINT6464 bit32vf.reg_tensor为类型声明不能直接调用由编译器在赋值形式中自动声明如reg vf.load_align(...)或reg vf.max(...)。寄存器在pypto_pro.language.vector_function函数内创建和使用函数结束后自动释放创建后必须通过vf.load_align或vf.full初始化否则内容未定义。RegTensor 寄存器数量上限为 32。超出上限的寄存器数据会写入预留的 8K UB 内存可能引起性能劣化编译器会自动复用生命周期结束的寄存器和预留内存若两者均有可用空间优先复用寄存器。运算开关mask_reg 掩码寄存器preg是掩码寄存器mask_reg是 VF 计算的元素级有效性控制容器详见 mask_reg.md总位宽固定为 256 bit粒度由关联的 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 bitVF 算子执行时根据 mask_reg 中每个数据元素对应的比特位决定该元素是否参与运算比特位为 1有效该元素参与运算结果写入目的寄存器对应位置比特位为 0无效该元素不参与运算目的寄存器对应位置置零vf.full等少数算子支持通过 mode 参数选择保留原值。mask_reg 同样不能直接调用由编译器在赋值形式中自动声明如preg vf.create_mask(...)或preg vf.eq(...)其 dtype 一般与配套的 reg_tensor 一致不一致时需要自行判断结果行为。MaskReg 寄存器数量上限为 16。未选中元素的处理MergeModemode参数对应 MergeMode 枚举定义 VF 计算指令中 mask 未选中元素非活跃元素在目标寄存器中的处理方式class MergeMode(enum.Enum): ZEROING ... # mask未选中位置置零默认 MERGING ... # mask未选中位置保留目标寄存器原值pypto_pro.language.MergeMode.ZEROING默认preg 未筛选的元素在 dst 中置 0pypto_pro.language.MergeMode.MERGING当前不支持。使用示例来自 MergeMode.mdpl.vector_function def vf_kernel(): dst vf.add(src0, src1, preg, modepl.MergeMode.ZEROING)vf.max的参数声明mode: Optional[MergeMode] None说明该参数可省略省略时使用默认的 ZEROING 语义与源码注释mode:pl.MergeMode.ZEROING(default). MERGING mode is not supported on current device.完全一致。约束说明符号零语义输入 src0 为 -0、src1 为 0 的情况下输出 dst 为 0。即max(-0.0, 0.0) 0.0与 IEEE 浮点规范中fmax的行为一致但需与 PyTorch 的torch.maximum对照验证见下文测试说明。数据类型一致性src0、src1、dst 数据类型需保持一致掩码粒度由 dtype 决定。不支持 FP8/FP4 直接运算虽然 reg_tensor 支持 FP8/FP4 存储类型但它们仅支持数据搬运load_align/store_align、数据填充full和类型转换astype不支持直接参与算术运算。vf.max支持的 dtype 列表DT_INT8 至 DT_UINT64中也不包含 FP8/FP4 类型需要时须先通过vf.astype转换为 FP32/BF16/FP16 再参与计算。返回值说明返回dst目的操作数类型为 reg_tensor支持的数据类型与 src0 中的说明一致。调用示例基本调用示例FP32以下示例完整演示了从全局张量经 UB Tile 加载到寄存器、执行vf.max、再写回并校验的完整流程见 max.mdimport 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(patternpl.MaskPattern.ALL, dtypepl.DT_FP32) reg_a vf.load_align(src_a, 0) reg_b vf.load_align(src_b, 0) reg_out vf.max(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_FP32], b: 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() in_b_grp pl.make_tile_group(typetf, addrs0x100, mutex_ids[1]) in_b in_b_grp.current() t_out_grp pl.make_tile_group(typetf, addrs0x200, 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 fnpu:{device_id} core_nums 1 torch.npu.set_device(device) a torch.randn([1, 64], devicedevice, dtypetorch.float32) b 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, torch.maximum(a, b), rtol1e-5, atol1e-5) if __name__ __main__: test_example() print(PASSED)代码结构拆解pl.vector_function定义 VF 计算函数example_vf接收 UB Tilesrc_a、src_b与目的 Tiledst_tile。函数内完成建掩码 → 加载寄存器 → 寄存器运算 → 存储结果四步vf.create_mask(patternpl.MaskPattern.ALL, dtypepl.DT_FP32)创建全 1 掩码FP32 粒度为 b3264 个元素全部有效vf.load_align(src_a, 0)从 UB Tile 偏移 0 处对齐加载数据到寄存器对应vlds指令见 python/pypto_pro/language/_vf_api.py 中load_align的注释Load aligned data from a UB Tile into a VF register (vlds instruction)vf.max(reg_a, reg_b, preg)逐元素求最大值未指定 mode使用默认 ZEROINGvf.store_align(dst_tile, reg_out, preg)将结果寄存器写回 UB Tile。pl.jit()定义宿主内核example_kernel声明三个pl.Tensor输入输出通过pl.TileType(shape[1, 64], dtypepl.DT_FP32, target_memorypl.MemorySpace.Vec)定义 UB Tile 形状——[1, 64]恰好对应 FP32 寄存器 64 个元素target_memorypl.MemorySpace.Vec表示向量内存空间UB。make_tile_group以不同addrs与mutex_ids为三个 Tile 分配互斥的 UB 地址0x0、0x100、0x200pl.section_vector()划定向量指令段。测试校验test_example通过TILE_FWK_DEVICE_ID环境变量默认 0选择 NPU 设备用torch.randn生成随机 FP32 数据example_kernelNone, core_nums以 1 个核启动内核最后用torch.testing.assert_close(out, torch.maximum(a, b), rtol1e-5, atol1e-5)将结果与 PyTorch 的torch.maximum对照验证正确性。说明vf.max仅支持在pl.vector_function函数内调用宿主pl.jit函数中不能直接使用。仓库测试 python/tests/st/pypto_pro/frontend/vf_api/test_vf_basic_ops.py 中的_vf_kernel_39_selectr_max_0展示了完全一致的模式vf.create_mask建掩码 →vf.load_align加载 →vf.max(reg_a, reg_b, preg)→vf.store_align写回可交叉印证接口的标准用法。INT64 数据类型示例vf.max支持 64 位整型此示例展示了 INT64 场景下 Tile 形状、地址对齐与掩码粒度的适配见 max.mdimport os import pypto_pro.language as pl import torch import torch_npu pl.vector_function def example_vf_int64(src_tile_a, src_tile_b, dst_tile): preg vf.create_mask(patternpl.MaskPattern.ALL, dtypepl.DT_INT64) reg_a vf.load_align(src_tile_a, 0) reg_b vf.load_align(src_tile_b, 0) reg_out vf.max(reg_a, reg_b, preg) vf.store_align(dst_tile, reg_out, preg) pl.jit() def example_kernel_int64( a: pl.Tensor[[pl.DYNAMIC, pl.DYNAMIC], pl.DT_INT64], b: pl.Tensor[[pl.DYNAMIC, pl.DYNAMIC], pl.DT_INT64], out: pl.Tensor[[pl.DYNAMIC, pl.DYNAMIC], pl.DT_INT64], ): tf pl.TileType(shape[1, 32], dtypepl.DT_INT64, target_memorypl.MemorySpace.Vec) in_a_grp pl.make_tile_group(typetf, addrs0, mutex_ids[0]) in_a in_a_grp.current() in_b_grp pl.make_tile_group(typetf, addrs256, mutex_ids[1]) in_b in_b_grp.current() t_out_grp pl.make_tile_group(typetf, addrs512, 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_int64(in_a, in_b, t_out) pl.store(out, t_out, [0, 0]) def test_example_int64(): 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.randint(-100, 100, [1, 32], devicedevice, dtypetorch.int64) b torch.randint(-100, 100, [1, 32], devicedevice, dtypetorch.int64) out torch.empty([1, 32], devicedevice, dtypetorch.int64) example_kernel_int64None, core_nums torch.npu.synchronize() torch.testing.assert_close(out, torch.maximum(a, b), rtol0, atol0) if __name__ __main__: test_example_int64() print(PASSED)与 FP32 示例的差异要点Tile 形状INT64 寄存器固定 256 字节只能容纳 32 个元素因此pl.TileType(shape[1, 32], ...)UB 地址三个 Tile 的地址分别为 0、256、512。每个 INT64 Tile 占 32 × 8 256 字节地址按 256 字节递增恰好与寄存器总大小对齐掩码粒度INT64 的 mask 粒度为 b64每元素 8 bitcreate_mask(patternpl.MaskPattern.ALL, dtypepl.DT_INT64)将 32 个元素全部置为有效测试数据使用torch.randint(-100, 100, ...)生成整数数据校验容差为rtol0, atol0整型精确比对。底层实现与源码印证接口声明与语义一致性vf.max的 docstring 与官方文档一一对应python/pypto_pro/language/_vf_api.pyFor each laneiwheremask[i]is active, compares the corresponding elements ofsrc0andsrc1and writes the larger value todst[i].——对应功能说明中的逐元素求最大值语义数学公式dstReg_i max(srcReg0_i, srcReg1_i)——与文档中的dst_i max(src0_i, src1_i)完全一致mode:pl.MergeMode.ZEROING(default). MERGING mode is not supported on current device.——对应参数说明中的 mode 约束。掩码驱动的 lane 级执行模型从 mask_reg.md 与 reg_tensor.md 可以看出整个 VF 计算模型以lane为单位每个 lane 对应一个数据元素其有效性由 mask_reg 的对应比特位决定。vf.max执行时只对活跃 lane 做比较非活跃 lane 按 MergeMode默认 ZEROING置零。这一模型同样适用于vf.add、vf.sub、vf.mul、vf.div等二元计算接口MergeMode 文档中明确列举了这些接口。赋值形式的编译器自动声明VF 运算接口含vf.max不返回可手动创建的目的寄存器而是通过赋值形式由编译器自动声明。从解析器源码 python/pypto_pro/language/parser/_assignment_parser.py 的注释可以看到Intercept VF op assignment form: reg vf.xxx(...) or reg_lo, reg_hi vf.xxx(...)——赋值形式拦截机制Only for compute ops that have dst registers (dst_count 0). Declaration ops (create_mask, RegTensor, compare, etc.) use normal return-value assignment and must NOT be intercepted.——vf.max属于具有目的寄存器的计算算子dst_count 0走拦截路径自动声明 reg_tensor而vf.create_mask属于声明算子走普通返回值赋值路径。同时 python/pypto_pro/language/parser/_call_parser.py 定义了_VF_MASK_DST_OPSeq、ne、lt 等产生 MaskReg 的算子集合与_VF_MASK_PRODUCING_OPScreate_mask、update_mask、get_mask_spr、mask_gen_with_reg_tensor 等产生掩码寄存器的算子集合编译器据此推断统一算子unified op的目的寄存器种类。这也解释了为什么示例中必须先通过vf.create_mask显式创建preg再将其作为vf.max的第三个参数传入。寄存器资源上限与性能提示reg_tensor上限 32、mask_reg上限 16当 VF 函数内同时存活的寄存器数量超过上限时编译器会将多余寄存器数据溢出到预留的 8K UB 内存并自动复用生命周期结束的寄存器这可能导致性能劣化。因此在编写包含多次vf.max的复杂 VF 函数时应尽量复用寄存器变量、缩短寄存器生命周期避免不必要的溢出。与相关接口的区分与配合vf.max二元寄存器版 vsvf.maxs寄存器-标量版vf.maxs(reg_a, scalar, preg)是寄存器与标量常量的逐元素最大值运算而vf.max是两个寄存器之间的运算。仓库测试 test_vf_basic_ops.py 中_vf_kernel_7_maxs_lrelu_0展示了vf.maxs(reg_a, -0.5, preg)的用法配合vf.leaky_relu实现 LeakyReLU 激活。实际算子实现中max与maxs经常成对出现。vf.min/vf.minsvf.min为按元素求最小值签名与vf.max完全对称可用于 Clamp 类运算与max配合。pl.maximum张量级pl.maximum是作用于 Tile/Tensor 层级的逐元素最大值接口适用于不需要寄存器级精细控制的场景vf.max则用于需要在向量函数内以寄存器粒度控制、配合掩码做条件运算的场景。两者处于不同抽象层级。pl.max标量级pypto_pro.language.max是标量运算用于循环边界计算等见 python/pypto_pro/language/_api.py 的注释 Scalar-only operation for loop-bound calculations etc. For tile element-wise maximum, usepl.maximum.与vf.max的寄存器语义完全不同使用时注意区分命名空间。常见使用场景激活函数实现如 LeakyReLU 可用max(x, 0) min(x, 0) * slope组合实现vf.max与vf.min、vf.muls配合可在寄存器级完成数值裁剪Clampvf.max(vf.min(x, upper), lower)实现上下界裁剪注意力机制FlashAttention仓库中 python/tests/st/pypto_pro/frontend/fa/test_fa_with_mask.py、python/tests/st/pypto_pro/frontend/fa/test_flex_attention.py 等测试在 FA 的 online softmax 更新路径running max 更新中使用vf.max这是寄存器级 max 最典型的工业级应用场景掩码条件运算借助vf.create_mask生成的局部掩码MaskPattern.VL1..VL128、M3、M4、H、Q 等可只对指定 lane 子集执行最大值运算非选中元素默认置零。总结vf.max是 CANN PyPTO 向量函数体系中按元素求最大值的核心二元寄存器运算接口其掩码筛选 默认 ZEROING的执行模型贯穿整个 VF 计算范式。使用时需把握三个关键点一是通过vf.create_mask显式创建与数据 dtype 匹配的掩码寄存器二是遵循load_align → 运算 → store_align的寄存器数据流目的寄存器由编译器在赋值形式中自动声明三是注意数据类型一致性、寄存器数量上限RegTensor 32 / MaskReg 16以及 FP8/FP4 仅支持搬运与类型转换的约束。结合 max.md 中 FP32 与 INT64 两个完整示例以及仓库源码与测试用例即可在 Ascend 950 系列平台上快速落地可验证的寄存器级逐元素最大值运算。【免费下载链接】pyptoPyPTO发音: pai p-t-oParallel Tensor/Tile Operation编程范式。项目地址: https://gitcode.com/cann/pypto创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表