ARTICLE DETAIL

资讯详情

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

PyPTO SIMD-API 深度解析:DuplicatePos 枚举与 vf.full 寄存器广播位置控制

PyPTO SIMD-API 深度解析:DuplicatePos 枚举与 vf.full 寄存器广播位置控制 PyPTO SIMD-API 深度解析DuplicatePos 枚举与 vf.full 寄存器广播位置控制【免费下载链接】pyptoPyPTO发音: pai p-t-oParallel Tensor/Tile Operation编程范式。项目地址: https://gitcode.com/cann/pypto导读本文聚焦 CANN PyPTO 向量函数VFVector Function编程范式中的DuplicatePos枚举类型。它用于在vf.full的 Tensor 广播模式下精确指定将源reg_tensor中哪一个索引位置的元素复制到整个目标寄存器是控制寄存器级数据复制的关键语义开关。读完本文你将掌握DuplicatePos的原型定义、LOWEST/HIGHEST两个成员的行为差异、与vf.full各参数preg、mode、dtype的配合方式以及从 Python 前端到 IR、再到后端代码生成的完整调用链。产品支持情况DuplicatePos作为vf.full的pos参数类型其可用性随目标昇腾硬件产品而不同当前仓库文档明确的支持矩阵如下产品支持情况Ascend 950PR / Ascend 950DT支持Atlas A3 训练系列产品 / Atlas A3 推理系列产品不支持Atlas A2 训练系列产品 / Atlas A2 推理系列产品不支持也就是说DuplicatePos及其所属的 VF 寄存器广播能力目前仅在 Ascend 950 系列产品上可用。编写可移植内核时应针对目标产品做特性判断避免在 Atlas A2/A3 平台上调用依赖该枚举的vf.fullTensor 广播功能。功能说明广播位置选择器DuplicatePos定义了 vf.full 广播模式下元素的复制位置选择用于指定将哪个索引位置的元素广播到整个寄存器。在向量寄存器的数据搬运场景中广播broadcast意味着把一个元素复制到目标寄存器的所有 lane 上。当一个源reg_tensor含有多个元素时必须明确复制哪一个是索引最低的元素还是索引最高的元素。DuplicatePos正是这一决策的形式化表达它消除了广播操作中的歧义使内核语义显式化、可读化。原型定义DuplicatePos是基于enum.Enum的枚举类型完整原型如下class DuplicatePos(enum.Enum): LOWEST ... # 广播最低索引位置的元素默认 HIGHEST ... # 广播最高索引位置的元素两个枚举成员的含义LOWEST广播源reg_tensor中索引最低的元素。这是vf.full的默认行为不显式传pos参数时即采用该语义。HIGHEST广播源reg_tensor中索引最高的元素。当需要复制寄存器末位元素例如取浮点量化中的尾端元素、拼接结果中的末尾分量时使用。从仓库源码看该枚举在多个层次均有对应定义IR 属性层framework/include/ir/op_attr_types.h#L165 通过宏PYPTO_DECLARE_ENUM(DuplicatePos, LOWEST, HIGHEST)声明Python 绑定层python/src/bindings/ir/ir.cpp#L665-L672 使用 pybind11 的py::enum_ir::DuplicatePos导出为pypto.ir.DuplicatePos并附带文档字符串 Position selector for vf.full (Tensor/broadcast mode)语言命名空间在 python/pypto_pro/language/_vf_api.py#L31 与 python/pypto_pro/language/init.py#L211 中被导入pl命名空间内核中通过pl.DuplicatePos.LOWEST/pl.DuplicatePos.HIGHEST访问。使用场景vf.full 的两种广播模式DuplicatePos是 vf.full 的pos关键字参数而vf.full本身支持两种广播模式Scalar 模式将标量值广播到寄存器各元素对应底层vbr/vdup指令Tensor 模式将源reg_tensor的最低或最高位元素广播到目标reg_tensor各元素对应底层vdup指令此时必须携带掩码preg。需要强调的是DuplicatePos仅在 Tensor 模式下生效——Scalar 模式的源是标量不存在索引位置概念pos参数自然不适用。Tensor 模式中pos决定复制源寄存器的哪一端。参数全解与约束以下是vf.full的完整函数原型与参数说明源自 full.md与DuplicatePos密切相关的pos参数已重点标注full(src, preg, dtype: Optional[DType] None, mode: Optional[MergeMode] None, pos: Optional[DuplicatePos] None) - dst参数输入/输出说明src输入源操作数为标量值或者 reg_tensor。源操作数 src 与目的操作数 dst 的数据类型保持一致。-Scalar 模式标量值广播到寄存器各元素。支持的数据类型为DT_INT8、DT_UINT8、DT_INT16、DT_UINT16、DT_FP16、DT_BF16、DT_INT32、DT_UINT32、DT_FP32、DT_INT64、DT_UINT64。-Tensor 模式reg_tensor广播其最低位或最高位元素。支持的数据类型为DT_INT8、DT_UINT8、DT_INT16、DT_UINT16、DT_FP16、DT_BF16、DT_INT32、DT_UINT32、DT_FP32、DT_INT64、DT_UINT64、DT_FP8E4M3FN、DT_FP8E5M2、DT_FP8E8M0、DT_HF8、DT_FP4E2M1、DT_FP4E1M2。preg输入mask_reg。Tensor 模式必选Scalar 模式可选。dtype输入可选指定数据类型。Scalar 模式必须输入Tensor 模式可从源寄存器自动推断。mode输入可选对应 MergeMode 类型。- pypto_pro.language.MergeMode.ZEROING默认preg 未筛选的元素在 dst 中置 0。- pypto_pro.language.MergeMode.MERGING 当前不支持。pos输入可选Tensor 模式下选择广播源 reg_tensor 的哪个元素对应DuplicatePos类型- pypto_pro.language.DuplicatePos.LOWEST默认广播最低位的元素。- pypto_pro.language.DuplicatePos.HIGHEST指定广播最高位的元素。关键约束与要点Tensor 模式必须带掩码preg为必选项用于限定广播写入的 lane 范围未筛选元素的行为由mode决定默认 ZEROING 置 0。默认语义pos缺省时等价于pl.DuplicatePos.LOWEST广播最低位元素。MERGING 模式限制当前设备上MergeMode.MERGING不支持原因可追溯到后端实现——backend_cce_vf_ops.cpp#L895-L896 中注释明确指出 MERGING is not supported by the underlying vdup/vbr instructions on current device并调用VFZeroingOnly(op, vf.full)强制 ZEROING 语义。约束说明vf.full本身无其他约束。返回值dst为目的操作数 reg_tensor支持的数据类型与src一致。调用示例最小示例两种位置选择的对比import pypto_pro.language as pl pl.vector_function def vf_kernel(): dst vf.full(preg, pospl.DuplicatePos.LOWEST)完整示例HIGHEST 与 LOWEST 的端到端验证下面给出一个可在昇腾 950 环境上运行的完整示例取自仓库前端 API 声明与测试用例的组合在一个内核中分别演示HIGHEST与LOWEST的广播差异并写入两个不同的输出 Tileimport os import pypto_pro.language as pl import torch import torch_npu pl.vector_function def vf_dup_example(in_a, t_f0, t_f1): preg vf.create_mask(patternpl.MaskPattern.ALL, dtypepl.DT_FP32) reg_a vf.load_align(in_a, 0) # 广播源寄存器最高位元素 reg_highest vf.full(reg_a, preg, pospl.DuplicatePos.HIGHEST) vf.store_align(t_f0, reg_highest, preg) # 广播源寄存器最低位元素默认行为 reg_lowest vf.full(reg_a, preg, pospl.DuplicatePos.LOWEST) vf.store_align(t_f1, reg_lowest, preg) pl.jit() def example_kernel( a: pl.Tensor[[pl.DYNAMIC, pl.DYNAMIC], pl.DT_FP32], out_h: pl.Tensor[[pl.DYNAMIC, pl.DYNAMIC], pl.DT_FP32], out_l: pl.Tensor[[pl.DYNAMIC, pl.DYNAMIC], pl.DT_FP32], ): tf pl.TileType(shape[1, 64], dtypepl.DT_FP32, target_memorypl.MemorySpace.Vec) in_a pl.make_tile(tf, addr0x0) t_h pl.make_tile(tf, addr0x100) t_l pl.make_tile(tf, addr0x200) with pl.section_vector(): pl.load(in_a, a, [0, 0]) vf_dup_example(in_a, t_h, t_l) pl.store(out_h, t_h, [0, 0]) pl.store(out_l, t_l, [0, 0]) def test_example(): device_id int(os.environ.get(TILE_FWK_DEVICE_ID, 0)) device fnpu:{device_id} torch.npu.set_device(device) a torch.arange(64, devicedevice, dtypetorch.float32).reshape([1, 64]) out_h torch.empty([1, 64], devicedevice, dtypetorch.float32) out_l torch.empty([1, 64], devicedevice, dtypetorch.float32) example_kernelNone, 1 torch.npu.synchronize() # HIGHEST所有 lane 都等于源寄存器最后一个元素 a[0, 63] torch.testing.assert_close(out_h, torch.full([1, 64], 63.0, devicedevice), rtol1e-5, atol1e-5) # LOWEST所有 lane 都等于源寄存器第一个元素 a[0, 0] torch.testing.assert_close(out_l, torch.full([1, 64], 0.0, devicedevice), rtol1e-5, atol1e-5) if __name__ __main__: test_example() print(PASSED)与其他数据类型的组合使用vf.full的 Tensor 模式覆盖从 INT 到 FP 的丰富数据类型DuplicatePos在这些场景下行为一致。仓库 full.md 中提供了 FP8E4M3FN、FP4E1M2、HF8、FP8E8M0、INT64 等多种完整示例核心模式均为先用vf.load_align将数据加载到寄存器再用vf.full(reg, preg)可搭配pos做广播最后vf.store_align写回。例如 FP8E8M0 场景中源寄存器首元素被置为0x7E后广播到整个输出验证逻辑为expected torch.full([1, 256], 0x7E, ...)INT64 场景则验证vf.full(42, preg, dtypepl.DT_INT64)的 Scalar 广播。这些示例同时印证Tensor 模式下dtype可从源寄存器自动推断而 Scalar 模式必须显式传入dtype。源码级原理从参数到指令的完整链路Python 前端参数解析与枚举校验在 python/pypto_pro/language/parser/_call_parser.py#L918 中_VF_KWARG_ENUMS将关键字pos映射到(DuplicatePos,)枚举元组。这意味着内核源码被 AST 解析器捕获时pos参数会被严格校验为DuplicatePos类型从而在编译早期拦截非法取值。vf.full的 Python 侧声明位于 python/pypto_pro/language/_vf_api.py#L104-L133其 docstring 与仓库文档保持一致Scalar 模式对应vbr/vdup指令Tensor 模式对应vdup指令且注明 Tensor 模式必须带掩码。IR 层与后端区分标量/向量源广播后端代码生成逻辑位于 framework/src/interface/pypto_pro/backend/backend_cce_vf_ops.cpp#L895-L919其判定流程清晰地展示了DuplicatePos在编译期的作用读取pos关键字ir::EnumToString(static_castir::DuplicatePos(op-GetKwargint(pos)))将枚举值转为LOWEST/HIGHEST字符串判定向量源广播只要显式传入了pos即认定是向量源广播is_vector_src !pos.empty()若未传pos则进一步检查src是否为RegTensor变量Tensor 模式生成vdup(dst, src_vec, mask, POS_xxx, MODE)指令其中POS_xxx即由DuplicatePos转换而来LOWEST→POS_LOWESTHIGHEST→POS_HIGHEST。由此可以推断pos参数是编译期区分 Scalar 广播与 Tensor 广播的重要信号之一即使不显式传pos只要源是寄存器变量编译器也能识别为向量广播并采用默认的LOWEST语义。测试验证仓库测试 python/tests/st/pypto_pro/frontend/vf_api/test_vf_basic_ops.py#L2819-L2828 中_vf_kernel_79_dup_highest_hist_freq_0内核分别用pospl.DuplicatePos.HIGHEST与pospl.DuplicatePos.LOWEST调用vf.full并将结果存入两个不同的输出 Tile 以对比差异单元测试 framework/tests/ut/interface/src/pypto_pro/backend/test_backend_cce_vf_ops.cpp#L278 与 #L387 则直接以EnumValue(ir::DuplicatePos::HIGHEST)和EnumValue(ir::DuplicatePos::LOWEST)构造操作属性验证后端代码生成路径。这为读者提供了双重的行为参照既可从 ST 用例观察运行时语义也可从 UT 用例观察 IR/后端层面的处理。小结与实践建议DuplicatePos是 PyPTO VF 编程中一个语义清晰但作用关键的枚举类型默认值优先绝大多数场景下广播最低位元素LOWEST即为所需语义此时可省略pos参数显式优于隐式当广播源是拼接结果向量归约产物等元素含义不对称的寄存器时建议显式写出pospl.DuplicatePos.HIGHEST或LOWEST让内核意图一目了然注意平台差异该能力目前仅 Ascend 950PR/950DT 支持Atlas A2/A3 系列不支持跨平台移植时需评估降级方案结合掩码使用Tensor 模式必须携带preg未筛选 lane 默认置 0ZEROING需要保留原值语义时需自行设计掩码策略。如需进一步了解广播操作的整体参数体系可继续阅读 vf.full 完整文档、MergeMode、reg_tensor 与 mask_reg 的说明。【免费下载链接】pyptoPyPTO发音: pai p-t-oParallel Tensor/Tile Operation编程范式。项目地址: https://gitcode.com/cann/pypto创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表