拓冰建站拓冰建站
首页 / 资讯中心 / 正文

pto-isa TRSQRT 指令全解:逐元素倒数平方根的语义、内建接口与 NPU/CPU 实现剖析

pto-isa TRSQRT 指令全解逐元素倒数平方根的语义、内建接口与 NPU/CPU 实现剖析【免费下载链接】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本文基于 pto-isa 仓库的指令规格文档 TRSQRT_zh.md 展开完整覆盖 TRSQRTTile Reciprocal Square Root指令的数学语义、三种汇编语法形式、C 内建接口签名与约束条件、临时空间tmp的语义差异并结合 include/pto/npu/a5/TRsqrt.hpp 与 include/pto/cpu/TRSqrt.hpp 等仓库源码剖析该指令在 NPU 向量域与 CPU 参考实现中的真实执行路径、精度参数RsqrtAlgorithm的作用机制以及高/默认精度分支的编译期选择逻辑。读完本文你可以在 PTO kernel 中正确使用TRSQRT(dst, src)/TRSQRT(dst, src, tmp)内建接口理解其精度控制手段与静态/运行时约束并能对照汇编语法与向量指令实现做代码评审或性能分析。指令定位与数学语义TRSQRT 是 PTO 虚拟指令集Virtual ISA中的一条逐元素倒数平方根指令属于向量域的逐元素超越函数类操作。它对 Tile 有效区域valid region内的每个元素(i, j)执行$$ \mathrm{dst}{i,j} \frac{1}{\sqrt{\mathrm{src}{i,j}}} $$即dst 1 / sqrt(src)。这类指令常见于归一化如 LayerNorm / RMSNorm 中的1/sqrt(mean(x^2))、量化缩放因子计算等场景把先开方再求倒数的复合操作压缩为一条 Tile 级指令减少向量寄存器往返与调度开销。从源码结构看指令的调度归属可以追溯到 include/pto/common/event.hpp 中的Op::TRSQRT枚举并通过PTO_DEFINE_OP_PIPE(Op::TRSQRT, PIPE_V)第 192 行将其绑定到向量流水线PIPE_V印证了它是一条纯粹的向量域逐元素指令。汇编语法同步形式与两个抽象层AS LevelPTO 汇编为同一条指令提供了三层表达。以下语法形式均出自 TRSQRT_zh.md同步形式Synchronous form%dst trsqrt %src : !pto.tile...AS Level 1SSA 形式——显式声明输入类型与输出类型%dst pto.trsqrt %src : !pto.tile... - !pto.tile...AS Level 2DPS 形式——以ins/outs操作数列表描述数据流使用tile_buf类型pto.trsqrt ins(%src : !pto.tile_buf...) outs(%dst : !pto.tile_buf...)自动模式与手动模式的 ASM 示例自动模式Auto资源放置与调度由编译器/运行时负责只需描述数据流# 自动模式由编译器/运行时负责资源放置与调度。 %dst pto.trsqrt %src : !pto.tile... - !pto.tile...手动模式Manual先显式绑定 Tile 资源再发射指令。绑定通过pto.tassign完成Tile 地址即 Tile 缓冲区偏移# 手动模式先显式绑定资源再发射指令。 # 可选当该指令包含 tile 操作数时 # pto.tassign %arg0, tile(0x1000) # pto.tassign %arg1, tile(0x2000) %dst pto.trsqrt %src : !pto.tile... - !pto.tile...PTO 汇编的最终形式可以写为紧凑的两段组合%dst trsqrt %src : !pto.tile... # AS Level 2 (DPS) pto.trsqrt ins(%src : !pto.tile_buf...) outs(%dst : !pto.tile_buf...)手动模式下与 Tile 分配相关的资源绑定指令可参考 TASSIGN 指令文档。C 内建接口两个重载与精度模板参数C 内建接口声明于 include/pto/common/pto_instr.hpp公共包含头为pto/pto-inst.hpp。共有两个重载template typename TileDataDst, typename TileDataSrc, typename... WaitEvents PTO_INST RecordEvent TRSQRT(TileDataDst dst, TileDataSrc src, WaitEvents ... events); template typename TileDataDst, typename TileDataSrc, typename TileDataTmp, typename... WaitEvents PTO_INST RecordEvent TRSQRT(TileDataDst dst, TileDataSrc src, TileDataTmp tmp, WaitEvents ... events);对照 include/pto/common/pto_instr.hpp 的实现源码接口签名中还包含一个文档未展开、但对精度控制至关重要的非类型模板参数template auto PrecisionType RsqrtAlgorithm::DEFAULT, typename TileDataDst, typename TileDataSrc, typename... WaitEvents, std::enable_if_tall_events_vWaitEvents..., int 0 PTO_INST RecordEvent TRSQRT(TileDataDst dst, TileDataSrc src, WaitEvents... events) { detail::PtoWaitEvents(events...); TRSQRT_IMPLPrecisionType(dst, src); return {}; }关键细节PrecisionType模板参数默认值为RsqrtAlgorithm::DEFAULT。RsqrtAlgorithm枚举定义于 include/pto/common/type.hpp取值为DEFAULT与HIGH_PRECISION两种。调用方可以显式特化例如TRSQRTRsqrtAlgorithm::HIGH_PRECISION(dst, src)选择更高精度的求根/除法算法后文 NPU 实现一节会给出两种分支的向量指令级差异。WaitEvents可变参数通过std::enable_if_tall_events_vWaitEvents...做编译期约束确保变参只能是 PTO 事件对象函数体内detail::PtoWaitEvents(events...)在发射指令前先等待这些事件用于跨指令的依赖同步。返回值RecordEvent返回一个记录事件的句柄可用于后续指令的WaitEvents参数形成流水线依赖链。3 参重载的编译期约束is_tile_data_vTileDataTmp确保tmp参数必须是一个 Tile 数据类型。约束条件静态检查与运行时检查文档给出的约束在 NPU 实现中被逐条落实。对照 include/pto/npu/a5/TRsqrt.hpp 中TRSQRT_IMPL的实现约束分为编译期static_assert与运行时PTO_ASSERT两级约束级别说明TileData::DType必须是float或half编译期NPU 实现中同时接受别名类型float32_t/float16_t并额外要求DstTile::DType与SrcTile::DType完全一致std::is_same_vTile 位置必须是向量域Loc TileType::Vec编译期static_assert(DstTile::Loc TileType::Vec SrcTile::Loc TileType::Vec, ...)静态有效边界ValidRow Rows且ValidCol Cols编译期src 与 dst 各自独立检查Tile 布局必须是行主序isRowMajor编译期static_assert(DstTile::isRowMajor SrcTile::isRowMajor, TRSQRT: Not supported Layout type)运行时src.GetValidRow() dst.GetValidRow()且src.GetValidCol() dst.GetValidCol()运行时PTO_ASSERT(dstValidCol src.GetValidCol(), ...)有效区域语义该操作使用dst.GetValidRow()/dst.GetValidCol()作为迭代域越界元素不参与计算定义域 / NaN语义src 0或负数输入下的行为由目标平台定义target-definedPTO 层面不做统一承诺注意一个实现细节NPU 路径迭代域取自dst的运行时有效行/列dstValidRow dst.GetValidRow()再用PTO_ASSERT强制要求 src 与之一致这与文档中该操作使用 dst 的 GetValidRow()/GetValidCol() 作为迭代域的描述一致。临时空间tmp2 参重载与 3 参重载的语义差异无tmp2 参数重载TRSQRT(dst, src)不需要tmp。默认精度实现直接使用vsqrtvdiv完成计算见下文 NPU 实现分析中间结果只占用向量寄存器。带tmp3 参数重载TRSQRT(dst, src, tmp)tmp被接口接受但当前 Ascend 950PR / Ascend 950DT对应 A5 系列实现并不使用它。源码证据在 include/pto/npu/a5/TRsqrt.hpptemplate auto PrecisionType RsqrtAlgorithm::DEFAULT, typename DstTile, typename SrcTile, typename TmpTile PTO_INTERNAL void TRSQRT_IMPL(DstTile dst, SrcTile src, TmpTile tmp) { TRSQRT_IMPLPrecisionType(dst, src); }3 参重载只是简单委托给 2 参实现tmp形参被完全忽略。CPU 参考实现 include/pto/cpu/TRSqrt.hpp 中的 3 参版本同样是空壳委托template auto PrecisionType RsqrtAlgorithm::DEFAULT, typename TileDataDst, typename TileDataSrc, typename TmpTileData PTO_INTERNAL void TRSQRT_IMPL(TileDataDst dst, TileDataSrc src, TmpTileData tmp) { TRSQRT_IMPL(dst, src); }tmp保留在 C 内建接口签名中是为了API 兼容性和潜在的未来高精度路径。如果你的 kernel 目前不需要临时 Tile使用 2 参重载即可如果你为将来的高精度版本预留了tmpTile传上去不会引入额外内存占用或计算。NPU 实现剖析vsqrt vdiv与高精度分支NPUA5 系列的完整实现在 include/pto/npu/a5/TRsqrt.hpp。整个实现分为四个层次1. 核心计算循环向量域三个底层函数TRsqrt_1D_NoPostUpdate、TRsqrt_1D_PostUpdate、TRsqrt_2D共享同一段向量计算逻辑以 include/pto/npu/a5/TRsqrt.hpp 的 1D 无后更新版本为例__VEC_SCOPE__ { RegTensorT srcReg; RegTensorT dstReg; RegTensorT tmpReg; RegTensorT oneReg; unsigned sReg validRow * validCol; MaskReg pReg CreatePredicateT(tmp); vdup(oneReg, (T)1.0, pReg, MODE_MERGING); for (uint16_t i 0; i repeatTimes; i) { pReg CreatePredicateT(sReg); vlds(srcReg, src, i * nRepeatElem, NORM); if constexpr (std::is_same_vT, float PrecisionType RsqrtAlgorithm::HIGH_PRECISION) { SqrtFloatImplT, RegTensorT(tmpReg, srcReg, pReg); DivIEEE754FloatImplT, RegTensorT(dstReg, oneReg, tmpReg, pReg); } else if constexpr (std::is_same_vT, half PrecisionType RsqrtAlgorithm::HIGH_PRECISION) { SqrtPrecisionImplT, RegTensorT(tmpReg, srcReg, pReg); DivIEEE754HalfImplT, RegTensorT(dstReg, oneReg, tmpReg, pReg); } else { vsqrt(tmpReg, srcReg, pReg, MODE_ZEROING); vdiv(dstReg, oneReg, tmpReg, pReg); } vsts(dstReg, dst, i * nRepeatElem, distValue, pReg); } }这段代码揭示了三个关键实现事实vsqrtvdiv组合默认精度路径先在tmpReg中算sqrt(src)再执行1.0 / sqrt_result与文档默认精度实现直接使用 vsqrt vdiv的描述完全一致。tmpReg是向量寄存器而非 Tile这解释了为什么 3 参重载中的tmpTile 用不上。精度分支是编译期选择if constexpr依据PrecisionType模板参数与元素类型float/half在编译期裁剪分支。HIGH_PRECISION路径调用SqrtFloatImpl/SqrtPrecisionImpl求根实现于 include/pto/npu/a5/custom/TSqrtHp.hpp与DivIEEE754FloatImpl/DivIEEE754HalfImplIEEE 754 精确除法实现于 include/pto/npu/a5/custom/Div754.hpp默认路径则是硬件向量指令vsqrtMODE_ZEROING加vdiv。谓词掩码保证越界安全CreatePredicateT(sReg)生成的MaskReg保证只有有效区域内的元素参与 load/store无效区域被屏蔽这正是有效区域语义的落地方式。nRepeatElem CCE_VL / sizeof(T)每次向量操作处理的元素数由 CCECube Engine / Vector 单元的向量长度决定repeatTimes CeilDivision(validRow * validCol, nRepeatElem)决定循环次数。2. 1D / 2D 调度与 Post-Update 优化顶层入口TRsqrtinclude/pto/npu/a5/TRsqrt.hpp根据 Tile 形状选择执行形态if constexpr ( ((DstTile::ValidCol DstTile::Cols) (SrcTile::ValidCol SrcTile::Cols)) || ((DstTile::Rows 1) (SrcTile::Rows 1))) { TRsqrt_1D_SwitchOp, PrecisionType, T, DstTile, SrcTile, nRepeatElem(dst, src, validRow, validCol, version); } else { TRsqrt_2DOp, PrecisionType, T, DstTile::RowStride, SrcTile::RowStride, nRepeatElem( dst, src, validRow, validCol); }当有效列等于物理列整行无边界缺口或 Tile 为单行时走1D 展平路径把validRow * validCol视为连续一维缓冲循环更简单、开销更低否则走2D 路径TRsqrt_2D逐行按SrcRowStride/DstRowStride寻址正确处理行步长不连续的情况1D 路径内部再由TRsqrt_1D_Switch依据VFImplKindVFIMPL_1D_NO_POST_UPDATE/VFIMPL_1D_POST_UPDATE/VFIMPL_2D_*选择是否使用POST_UPDATE寻址模式——POST_UPDATE版本把指针递增折叠进vlds/vsts指令本身如vlds(srcReg, src, nRepeatElem, NORM, POST_UPDATE)减少每次迭代的地址计算。3. 入口静态断言与运行时断言前文约束条件一节列出的静态断言均集中出现在TRSQRT_IMPLDstTile, SrcTileinclude/pto/npu/a5/TRsqrt.hpp中与文档约束逐条对应且额外补充了两条文档未单列的检查src/dst 数据类型必须完全一致以及 DType 对float32_t/float16_t别名的等价接受。CPU 参考实现CPU 端参考实现位于 include/pto/cpu/TRSqrt.hpp逻辑直观便于理解指令的语义基准template auto PrecisionType RsqrtAlgorithm::DEFAULT, typename TileDataDst, typename TileDataSrc PTO_INTERNAL void TRSQRT_IMPL(TileDataDst dst, TileDataSrc src) { static_assert(/* DType 均为 half 或均为 float */, TRSQRT: Invalid data type); static_assert( TileDataSrc::ValidRow TileDataDst::ValidRow TileDataSrc::ValidCol TileDataDst::ValidCol, TRSQRT: Src valid row/col ! Dst valid row/col); const std::size_t rows static_caststd::size_t(dst.GetValidRow()); const std::size_t cols static_caststd::size_t(dst.GetValidCol()); if (rows 0 || cols 0) { return; } cpu::parallel_for_rows(rows, cols, { for (std::size_t c 0; c cols; c) { const auto x static_castdouble(src.data()[GetTileElementOffsetTileDataSrc(r, c)]); const double y 1.0 / std::sqrt(x); dst.data()[GetTileElementOffsetTileDataDst(r, c)] static_casttypename TileDataDst::DType(y); } }); }CPU 实现的要点以 double 精度计算1.0 / std::sqrt(x)先升精度到double计算再截断回DType作为数值参考基准精度高于 NPU 默认路径的硬件近似按行并行cpu::parallel_for_rows将行维度分片并行行内列循环串行空有效区域快速返回rows 0 || cols 0时直接返回与 NPU 侧迭代域为 0的行为一致偏移量通过GetTileElementOffset定义于 include/pto/cpu/tile_offsets.hpp计算支持 Tile 的行步长布局。CPU 实现与 NPU 实现共同构成了测试体系见下文测试与验证中跨平台结果比对的基础。完整使用示例以下示例完整继承自 TRSQRT_zh.md展示了自动Auto与手动Manual两种编程模式下的用法。自动模式Auto#include pto/pto-inst.hpp using namespace pto; void example_auto() { using TileT TileTileType::Vec, float, 16, 16; TileT src, dst; TRSQRT(dst, src); }自动模式下Tile 的分配、地址绑定与调度由编译器/运行时统一管理开发者只描述数据流。手动模式Manual#include pto/pto-inst.hpp using namespace pto; void example_manual() { using TileT TileTileType::Vec, float, 16, 16; TileT src, dst; TASSIGN(src, 0x1000); // 将 src Tile 绑定到缓冲区偏移 0x1000 TASSIGN(dst, 0x2000); // 将 dst Tile 绑定到缓冲区偏移 0x2000 TRSQRT(dst, src); }手动模式下先用TASSIGN显式绑定 Tile 资源与汇编层的pto.tassign %arg, tile(addr)一一对应再发射TRSQRT指令。两个 Tile 均为TileTileType::Vec, float, 16, 16向量域、float类型、物理 16×16 元素。选择高精度路径如需更高精度可在调用时显式指定模板参数TRSQRTRsqrtAlgorithm::HIGH_PRECISION(dst, src);在 A5 目标上这会切换到 IEEE 754 精确求根/除法实现SqrtFloatImplDivIEEE754FloatImpl适用于对数值精度敏感的归一化类计算。测试与验证仓库的测试目录按平台组织了 TRSQRT 相关的测试用例可以作为行为验证与精度比对的入口tests/cpu/stCPU 参考实现的指令级测试tests/npu/a5A5 系列 NPU 目标上的 TRSQRT 测试含向量指令实现路径覆盖。这些测试与 CPU 参考实现double精度基准配合可用于验证 NPU 默认路径vsqrt vdiv与高精度路径在各自有效区域内的数值行为。相关指令与延伸阅读平方根指令 TSQRT与 TRSQRT 同族的超越函数指令接口形式类似同样带有SqrtAlgorithm精度模板参数Tile 资源绑定指令 TASSIGN手动模式下绑定 Tile 资源的基础指令向量域 Tile 类型与RsqrtAlgorithm等算法枚举include/pto/common/type.hpp事件与流水线映射Op::TRSQRT→PIPE_Vinclude/pto/common/event.hpp指令规格总入口docs/isa/README_zh.md。小结TRSQRT 是 pto-isa 中实现清晰、约束明确的逐元素倒数平方根指令数学语义为dst 1/sqrt(src)汇编层提供同步、AS Level 1SSA、AS Level 2DPS三种表达C 内建接口通过RsqrtAlgorithm非类型模板参数提供DEFAULT/HIGH_PRECISION两条编译期精度路径WaitEvents变参与RecordEvent返回值构成事件同步机制NPU 实现以vsqrt vdiv为默认路径、以 IEEE 754 精确求根/除法为高精度路径并依据 Tile 形状在 1D / 2D 与 Post-Update 寻址之间自动切换3 参重载中的tmpTile 当前未被使用仅为 API 兼容与未来高精度路径预留。理解这套接口与实现细节是正确、高效地在其上层编写归一化等数值敏感 kernel 的前提。【免费下载链接】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),仅供参考
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门