ARTICLE DETAIL

资讯详情

深耕商务建站与企业官网运营的一线实战洞察。

PTO-ISA TAXPY 指令深度解析:Tile 级原位缩放累加(a·x+y)的实现原理与实战

PTO-ISA TAXPY 指令深度解析:Tile 级原位缩放累加(a·x+y)的实现原理与实战 PTO-ISA TAXPY 指令深度解析Tile 级原位缩放累加a·xy的实现原理与实战【免费下载链接】pto-isaParallel Tile Operation (PTO) is a virtual instruction set architecture designed by Ascend CANN, focusing on tile-level operations. This repository offers high-performance, cross-platform tile operations across Ascend platforms.项目地址: https://gitcode.com/cann/pto-isa导读本文以 CANN pto-isa 仓库中的 docs/isa/TAXPY.md 为核心完整讲解 PTOParallel Tile Operation虚拟指令集中的TAXPY指令——它在 Tile 上执行原位缩放累加AXPY$a \cdot x y$。你将掌握 TAXPY 的数学语义、C 内建接口签名与参数约定、支持的数据类型组合、在向量流水线PIPE_V上的底层实现差异A5 与 A2/A3 两个后端路径、CPU_SIM 参考实现以及如何编写可在 NPU 上运行的完整 ST 测试内核。读完本文你可以在自己的 kernel 中安全、正确地使用TAXPY完成带标量缩放的累加运算。TAXPY 指令概述TAXPY对单个 Tile 执行原位缩放累加将源 Tilesrc0$x$按标量scalar$a$缩放后累加到目标 Tiledst$y$上结果写回dst$$ \mathrm{dst}{i,j} \leftarrow \mathrm{scalar} \cdot \mathrm{src0}{i,j} \mathrm{dst}_{i,j} $$三个操作数的角色非常明确见 docs/isa/TAXPY.mddst既是累加输入$y$也是输出调用前必须已初始化src0只读$x$逐元素参与运算scalar标量缩放系数$a$类型为TileDataSrc::DType。这与普通的TADDSTile 加标量不同TAXPY是两个 Tile 之间带标量系数的融合运算输入输出共用一个 TileRMW读-修改-写天然适合在已有累加基上做带权累加。数学语义基于有效区域的逐元素 RMW指令语义定义在 Tile 的有效区域valid region内。对有效区域中的每个元素(i, j)$$ \mathrm{dst}{i,j}^{\text{new}} \mathrm{scalar} \cdot \mathrm{src0}{i,j} \mathrm{dst}_{i,j}^{\text{old}} $$各操作数的数据流角色操作数角色说明dst读-修改-写RMW读入旧值作为累加基 $y$写回 $\mathrm{scalar} \cdot x y$src0只读逐元素贡献被缩放的 $x$scalar标量缩放系数 $a$类型为TileDataSrc::DType除非另有说明语义在有效区域内定义目标相关行为标记为实现定义implementation-defined。这一语义决定了两个关键约束dst必须先初始化否则累加基是未定义数据dst与src0的有效形状必须一一对应逐元素映射。C 内建接口签名、声明位置与参数约定TAXPY的公共声明位于 include/pto/common/pto_instr.hpp对外包含头为pto/pto-inst.hpp内部实现声明在pto/common/pto_instr.hpptemplate typename TileDataDst, typename TileDataSrc, typename... WaitEvents PTO_INST RecordEvent TAXPY(TileDataDst dst, TileDataSrc src0, typename TileDataSrc::DType scalar, WaitEvents ...events);参数方向含义dst输入/输出累加基与结果 Tile$y$读-修改-写Vecsrc0输入缩放源 Tile$x$只读Vec有效形状与dst相同scalar输入标量缩放系数$a$类型为TileDataSrc::DTypeevents...输入等待事件WaitEvents指令前隐式 event synchronization从源码看公开包装函数先调用detail::PtoWaitEvents(events...)完成隐式事件同步再经MAP_INSTR_IMPL(TAXPY, dst, src0, scalar)宏分发到各后端的TAXPY_IMPL实现参见 include/pto/common/pto_instr.hpp 的宏定义。返回值类型为RecordEvent可用于指令间依赖编排。Tile 尺寸与数据类型对于有效形状 $M \times N$ 的 TileTiledtype有效形状TileType说明dsthalf或float$M \times N$Vec(UB)累加基 结果RMWsrc0half或float$M \times N$Vec(UB)缩放源逐元素dst与src0的有效行数、有效列数必须完全相同。支持的 dtype 组合dstdtypesrc0dtypescalardtype说明halfhalfhalf同类型路径直接vaxpyfloatfloatfloat同类型路径直接vaxpyfloathalfhalf差异路径src0拓宽为 FP32 后累加dst与src0必须 dtype 一致或dst为float且src0为half允许 half→float 的拓宽累加。dst为half而src0为float的组合非法由实现内static_assert在编译期拦截。这一 dtype 规则并非只在文档中声明而是被真实编码进源码。在 include/pto/npu/a5/TAxpy.hpp 与 include/pto/npu/a2a3/TAxpy.hpp 的TAXPY_IMPL中均有两条static_assertstatic_assert(std::is_same_vT, half || std::is_same_vT, float, TAXPY: Invalid data type); static_assert(std::is_same_vT, U || (std::is_same_vT, float std::is_same_vU, half), TAXPY: The data type of dst must be consistent with src or dst is float while src is half.);同时static_assert(TileDataDst::Loc TileType::Vec, ...)保证两个 Tile 都位于 UB统一缓冲区、向量流水线并附带运行期PTO_ASSERT检查src0与dst的 valid row/col 一致。底层实现原理向量流水线上的 vaxpyTAXPY在向量流水线PIPE_V上执行核心是vaxpy$a \cdot x y$向量内建。不同后端在实现路径上有明显差异这正是 PTO 跨平台设计的关键体现。同类型路径A5 后端在 include/pto/npu/a5/TAxpy.hpp 的AxpyInstrSame中按CeilDivision(validCol, elementsPerRepeat)计算 repeat 次数elementsPerRepeat CCE_VL / sizeof(T)然后逐行、逐 repeatvlds加载src0与dstNORM 模式用CreatePredicateT(sreg)构造尾部谓词掩码屏蔽不足一个 repeat 的列执行vaxpy(vreg2, vreg0, scalar, preg)vsts写回dst。差异类型路径A5 后端UNPK_B16 vcvt在 include/pto/npu/a5/TAxpy.hpp 的AxpyInstrDiff中当dst为float、src0为half时vlds(..., UNPK_B16)以半精度解包模式加载src0vcvt(reg_src_tmp, vreg0, preg, PART_EVEN)将 half 拓宽为 FP32再以(T)scalar执行vaxpy并写回。即 A5 上差异路径需要显式“解包 类型转换 缩放累加”三步这与文档中“src0 拓宽为 FP32 后累加”的描述完全对应。A2/A3 后端count 模式与 norm 模式自适应在 include/pto/npu/a2a3/TAxpy.hpp 中A2/A3 的实现先计算常量dstStride / dstBlockSizeElem 255 || srcStride / srcBlockSizeElem 255判定repeat-stride 是否溢出uint8_t上限 255validCol / elementsPerRepeat validRow判定列 repeat 数与行数的关系。随后二选一count 模式AxpyCountModeinclude/pto/npu/a2a3/TAxpy.hppset_mask_count()SetVectorCount(validCol)逐行执行一次vaxpy(dstPtr, src0Ptr, scalar, 0, 1, 1, 8, 4)——注意差异类型下 src 一个 repeat 只占 4 个 blockdst 占 8 个 block由vaxpy原生处理 half→float 的 4-block src / 8-block dst 映射norm 模式AxpyNormMode/AxpyNormModeTailinclude/pto/npu/a2a3/TAxpy.hpp先按REPEAT_MAX切分行块行内按dstElementsPerRepeat循环切列剩余列用SetContMaskByDTypeU(numRemainAfterLoop)连续掩码处理尾部。选择逻辑useCountMode repeatStrideOverflow || validCol / elementsPerRepeat validRow保证任意有效形状包括大行数、超长列、窄行都能被覆盖这与文档中“按 repeat-stride 是否溢出、以及列数与行数的关系在 count 模式与 norm 模式间选择”的描述一一对应。CPU_SIM 参考实现模拟半精度舍入在 include/pto/cpu/TBinSOps.hpp 的 CPU_SIMTAXPY_IMPL中half/half路径刻意模拟硬件上的两步半精度行为volatile auto product static_casttypename TileDataDst::DType(src.data()[srcIdx] * scalar); volatile auto result static_casttypename TileDataDst::DType(dst.data()[dstIdx] product); dst.data()[dstIdx] result;即先把乘积舍入为half再与dst相加并把和再次舍入为half而不是按主机浮点精度一次性计算完整表达式。volatile保证两次舍入不被编译器优化合并从而与 NPU 上的可观测行为保持一致。float路径则直接按dst src * scalar计算。约束一览约束适用范围原因dst、src0必须为TileType::Vec所有目标在 UB向量流水线上执行dst与src0有效形状相同$M \times N$所有目标逐元素一一对应dstdtype ∈ {half,float}所有目标vaxpy支持的浮点字宽dst/src0dtype 一致或 (float,half)所有目标仅允许 half→float 拓宽累加dst调用前必须已初始化所有目标dst作为累加基 $y$ 被读入完整实战示例NPU 上的 ST 内核仓库在 tests/npu/a5/src/st/testcase/taxpy/taxpy_kernel.cpp 提供了完整的 ST 测试内核展示了TAXPY在真实 kernel 中的完整用法#include pto/pto-inst.hpp #include pto/common/constants.hpp #include acl/acl.h using namespace pto; template typename T, int kTRows_, int kTCols_, int vRows, int vCols __global__ AICORE void runTAxpy(__gm__ T __out__* out, __gm__ T __in__* src0, float scalar) { using DynShapeDim5 Shape1, 1, 1, vRows, vCols; using DynStridDim5 pto::Stride1, 1, 1, vCols, 1; using GlobalData GlobalTensorT, DynShapeDim5, DynStridDim5; using TileData TileTileType::Vec, T, kTRows_, kTCols_, BLayout::RowMajor, -1, -1; TileData src0Tile(vRows, vCols); TileData dstTile(vRows, vCols); TASSIGN(src0Tile, 0x0); TASSIGN(dstTile, 0x20000); GlobalData src0Global(src0); GlobalData dstGlobal(out); TLOAD(src0Tile, src0Global); TLOAD(dstTile, dstGlobal); set_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); TAXPY(dstTile, src0Tile, (T)scalar); // dst scalar * src0 dst set_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); TSTORE(dstGlobal, dstTile); out dstGlobal.data(); }要点解读dst必须先用TLOAD从 GM 载入作为累加基 $y$这与“调用前必须初始化”的约束一致TASSIGN为两个 Tile 指定 UB 上的起始地址set_flag/wait_flag是显式的 MTE2→V 流水线事件同步TAXPY自身的WaitEvents参数在这里未被用到但包装函数内部也会做隐式同步结果通过TSTORE写回 GM。该测试文件还覆盖了多种形状与 dtype 组合的实例化见 taxpy_kernel.cpp包括half的 64×64、63×63、1×16384、2048×16 以及float的 8×8、15×15 等从侧面验证了“任意有效形状均可覆盖”的实现目标。更多完整 ST 示例可参考tests/npu/a5/src/st/testcase/taxpy/A5tests/npu/a2a3/src/st/testcase/taxpy/A2/A3tests/npu/kirin9030/src/st/testcase/taxpy/Kirin9030tests/cpu/st/testcase/taxpy/CPU 参考实现相关资源指令文档docs/isa/TAXPY.md / docs/isa/TAXPY_zh.md公共声明与包装函数include/pto/common/pto_instr.hppA5 后端实现include/pto/npu/a5/TAxpy.hppA2/A3 后端实现include/pto/npu/a2a3/TAxpy.hppCPU_SIM 参考实现include/pto/cpu/TBinSOps.hpp公开包含头include/pto/pto-inst.hpp如需深入理解 Tile 编程模型与vaxpy之外的向量指令体系可继续阅读 docs/coding/ProgrammingModel.md 与 docs/isa/README.md。【免费下载链接】pto-isaParallel Tile Operation (PTO) is a virtual instruction set architecture designed by Ascend CANN, focusing on tile-level operations. This repository offers high-performance, cross-platform tile operations across Ascend platforms.项目地址: https://gitcode.com/cann/pto-isa创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表
PREV
查看更多资讯
NEXT
返回资讯列表