PTO 逐元素倒数平方根指令 TRSQRT 详解:从数学定义到 A5 向量内核实现
【免费下载链接】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
导读
TRSQRT 是 Ascend CANN 开源仓库 pto-isa 中基于 Parallel Tile Operation(PTO)虚拟指令集定义的逐元素倒数平方根(Reciprocal Square Root)tile 级运算指令,计算dst = 1 / sqrt(src)。本文以 TRSQRT 指令规范文档 为骨架,结合仓库内 A5、A2A3、CPU 多平台实现源码与 ST 测试用例,完整讲解其数学语义、两级汇编语法、C++ Intrinsic 接口、约束校验规则、精度路径选择与临时空间(tmp)的真实行为。读完本文,你将能够使用 C++ Intrinsic 或汇编形式编写 TRSQRT 算子,理解默认精度与高精度实现的内在差异,并知晓如何通过仓库测试用例验证结果正确性。
TRSQRT 的语义与数学定义
TRSQRT 是 PTO 指令集中的一种逐元素(Elementwise)一元运算,作用于向量(Vector)位置的 tile。它对有效区域(valid region)内的每个元素独立计算其倒数平方根。
对于有效区域内的每个元素(i, j),运算满足:
$$\mathrm{dst}{i,j} = \frac{1}{\sqrt{\mathrm{src}{i,j}}}$$
即先对输入取平方根,再求其倒数。该运算在归一化(如 RMSNorm、LayerNorm 中的1/sqrt(x)缩放项)、信号处理、数值算法中都是高频基础算子。TRSQRT 与仓库中的平方根指令 TSQRT 不同:TSQRT 只计算sqrt(src),而 TRSQRT 在其基础上增加一次倒数运算,文档明确其默认精度实现由vsqrt与vdiv两条底层向量指令组合完成。
需要特别说明的是,指令文档指出其定义域/NaN 行为是目标平台相关的(target-defined):例如当src == 0(除零)或输入为负数(负数的平方根无实数结果)时,不同硬件的具体表现可能不同,编写算子时应在上层做好输入约束。
汇编语法:两级抽象形式
TRSQRT 在 PTO 指令集中以三种形式出现,分别对应不同开发模式与汇编层级:
同步形式(PTO Assembly):
%dst = trsqrt %src : !pto.tile<...>AS Level 1(SSA 形式,用于 Auto 模式):
%dst = pto.trsqrt %src : !pto.tile<...> -> !pto.tile<...>AS Level 2(DPS 形式,显式绑定 tile buffer):
pto.trsqrt ins(%src : !pto.tile_buf<...>) outs(%dst : !pto.tile_buf<...>)其中!pto.tile<...>是 PTO 虚拟 ISA 中表示 tile 类型的方言类型,...处可展开为 tile 的形状、数据类型与布局等属性。Level 1(SSA)适合 Auto 模式——编译器/运行时负责 tile 的放置与调度;Level 2(DPS)则显式声明ins/outs的tile_buf,适用于需要精确控制资源绑定的手动开发场景。仓库 PTO-Virtual-ISA-Manual.md 对这两级汇编抽象有系统性描述。
C++ Intrinsic 接口
在应用开发中更常用的是 C++ Intrinsic。TRSQRT 在 include/pto/common/pto_instr.hpp 中声明了两个重载,均以RsqrtAlgorithm模板参数控制精度策略(默认RsqrtAlgorithm::DEFAULT):
template <auto PrecisionType = RsqrtAlgorithm::DEFAULT, typename TileDataDst, typename TileDataSrc, typename... WaitEvents, std::enable_if_t<all_events_v<WaitEvents...>, int> = 0> PTO_INST RecordEvent TRSQRT(TileDataDst &dst, TileDataSrc &src, WaitEvents &... events); template <auto PrecisionType = RsqrtAlgorithm::DEFAULT, typename TileDataDst, typename TileDataSrc, typename TileDataTmp, typename... WaitEvents, std::enable_if_t<is_tile_data_v<TileDataTmp> && all_events_v<WaitEvents...>, int> = 0> PTO_INST RecordEvent TRSQRT(TileDataDst &dst, TileDataSrc &src, TileDataTmp &tmp, WaitEvents &... events);两个重载的使用要点:
- 2 参数重载
TRSQRT(dst, src, events...):基本形式,通过TRSQRT_IMPL<PrecisionType>(dst, src)下发计算; - 3 参数重载
TRSQRT(dst, src, tmp, events...):多接收一个临时 tile,内部调用TRSQRT_IMPL<PrecisionType>(dst, src, tmp),为高精度路径预留接口; - 两者都返回
RecordEvent,可通过Event<Op::TRSQRT, Op::...>与前后指令(如 TLOAD、TSTORE)构成事件依赖链,实现异步流水调度; - 变参
WaitEvents受all_events_v<WaitEvents...>约束,确保只接受合法事件类型。
精度枚举RsqrtAlgorithm定义于 include/pto/common/type.hpp:
enum class RsqrtAlgorithm : uint8_t { DEFAULT, HIGH_PRECISION };约束与编译期校验
TRSQRT 对 tile 的形态有严格约束,这些约束在 NPU 实现中以static_assert(编译期)与PTO_ASSERT(运行期)双重方式执行。核心规则如下:
| 约束类别 | 规则 | 检查时机 |
|---|---|---|
| 数据类型 | DType必须为float或half | 编译期 |
| Tile 位置 | TileData::Loc == TileType::Vec,必须在向量单元 | 编译期 |
| 布局 | 必须为行主序TileData::isRowMajor | 编译期 |
| 静态有效边界 | ValidRow <= Rows且ValidCol <= Cols | 编译期 |
| 运行时形状匹配 | src.GetValidRow() == dst.GetValidRow()且src.GetValidCol() == dst.GetValidCol() | 运行期 |
| 迭代域 | 以dst.GetValidRow()/dst.GetValidCol()为迭代范围 | 运行期 |
以 include/pto/npu/a5/TRsqrt.hpp 中的TRSQRT_IMPL为例,可以看到一组完整的静态断言,任何一个不满足都会在编译阶段直接报错:
static_assert(DstTile::isRowMajor && SrcTile::isRowMajor, "TRSQRT: Not supported Layout type"); static_assert(DstTile::Loc == TileType::Vec && SrcTile::Loc == TileType::Vec, "TRSQRT: TileType of src and dst tiles must be TileType::Vec."); static_assert(DstTile::ValidCol <= DstTile::Cols, "TRSQRT: Number of dst's valid columns must not be greater than number of tile columns."); static_assert(DstTile::ValidRow <= DstTile::Rows, "TRSQRT: Number of dst's valid rows must not be greater than number of tile rows."); static_assert(std::is_same_v<typename DstTile::DType, typename SrcTile::DType>, "TRSQRT: The data type of dst must be consistent with of src"); static_assert(std::is_same_v<typename DstTile::DType, float32_t> || std::is_same_v<typename DstTile::DType, float> || std::is_same_v<typename DstTile::DType, float16_t> || std::is_same_v<typename DstTile::DType, half>, "TRSQRT: Invalid data type.");运行期则通过PTO_ASSERT检查 src/dst 的有效行列是否一致,随后读取dst.GetValidRow()与dst.GetValidCol()作为实际计算域。注意迭代域以dst的 valid 区域为准(文档明确说明),因此保证 src 与 dst valid 区域一致是正确性的前提。
临时空间(tmp)的真实行为:A5 上的接口兼容设计
这是 TRSQRT 一个容易被误用的设计点。指令文档明确说明了tmp参数在不同平台上的真实行为:
2 参数重载(无 tmp):不需要临时空间,默认精度实现直接用vsqrt+vdiv两条向量指令完成"先开方、再取倒数"。
3 参数重载(带 tmp):接口接收tmp,但当前 A5 实现并未使用它。从 include/pto/npu/a5/TRsqrt.hpp 可以看到,3 参数版本直接委托给 2 参数实现:
template <auto PrecisionType = RsqrtAlgorithm::DEFAULT, typename DstTile, typename SrcTile, typename TmpTile> PTO_INTERNAL void TRSQRT_IMPL(DstTile& dst, SrcTile& src, TmpTile& tmp) { TRSQRT_IMPL<PrecisionType>(dst, src); }tmp之所以保留在 C++ Intrinsic 签名中,是为 API 兼容性以及未来潜在的高精度路径预留。CPU 实现(include/pto/cpu/TRSqrt.hpp)同样采取委托策略。
值得注意的平台差异:A2A3 平台的实现并不相同。在 include/pto/npu/a2a3/TUnaryOp.hpp 中,3 参数重载会真正调用TRsqrtHighPrecision——先用vector_dup把tmp初始化为 1.0,再逐行执行vsqrt,经pipe_barrier(PIPE_V)同步后执行vdiv完成1/sqrt(x)。因此,如果为 A2A3 平台使用高精度路径,tmp是实际会被使用的临时缓冲;而 A5 上当前无论是否传入tmp,计算路径完全一致。
精度路径:DEFAULT 与 HIGH_PRECISION 的实现差异
从 include/pto/npu/a5/TRsqrt.hpp 的 1D 内核可以看出 A5 上两种精度策略的底层差异:
- DEFAULT(默认精度):对每个向量块执行
vsqrt(tmpReg, srcReg, pReg, MODE_ZEROING)后,再执行vdiv(dstReg, oneReg, tmpReg, pReg),其中oneReg通过vdup预加载为常量 1.0; - HIGH_PRECISION(高精度):对
float走SqrtFloatImpl+DivIEEE754FloatImpl,对half走SqrtPrecisionImpl+DivIEEE754HalfImpl,即使用符合 IEEE 754 语义的平方根与除法实现(相关实现见 include/pto/npu/a5/custom/TSqrtHp.hpp 与 include/pto/npu/a5/custom/Div754.hpp),以获得更精确的结果。
指令的循环次数由有效元素总数与单次 repeat 元素数决定:nRepeatElem = CCE_VL / sizeof(T)(向量处理单元一个 repeat 能处理的元素个数),repeat 次数为CeilDivision(validRow * validCol, nRepeatElem)。
A2A3 平台的默认路径则直接使用硬件近似倒数平方根指令vrsqrt(见 include/pto/npu/a2a3/TUnaryOp.hpp 中的RsqrtOp),高精度路径才改用vsqrt+vdiv组合。这体现了不同硬件代际在指令级实现上的差异。
使用示例:Auto 与 Manual 模式
Auto 模式
Auto 模式下 tile 的放置与调度由编译器/运行时管理,只需声明Tile<TileType::Vec, T, Rows, Cols>并直接调用 Intrinsic:
#include <pto/pto-inst.hpp> using namespace pto; void example_auto() { using TileT = Tile<TileType::Vec, float, 16, 16>; TileT src, dst; TRSQRT(dst, src); }Manual 模式
Manual 模式需要先用TASSIGN为 tile 显式绑定地址(UB 缓冲),再发起计算:
#include <pto/pto-inst.hpp> using namespace pto; void example_manual() { using TileT = Tile<TileType::Vec, float, 16, 16>; TileT src, dst; TASSIGN(src, 0x1000); TASSIGN(dst, 0x2000); TRSQRT(dst, src); }对应的汇编形式中,Manual 模式在指令前通过pto.tassign完成资源绑定(tile 操作数可省略,但显式绑定更清晰):
# Manual mode: resources must be bound explicitly before issuing the instruction. # Optional for tile operands: # pto.tassign %arg0, @tile(0x1000) # pto.tassign %arg1, @tile(0x2000) %dst = pto.trsqrt %src : !pto.tile<...> -> !pto.tile<...>Auto 模式则由编译器管理放置与调度,直接发射 SSA 形式的pto.trsqrt。
源码级内核实现解析(A5)
include/pto/npu/a5/TRsqrt.hpp 根据 tile 形态自动选择不同的内核变体。顶层分发逻辑TRsqrt先根据编译期信息判断:若ValidCol == Cols(整行有效)或行数为 1(纯一维数据),则走 1D 路径;否则走 2D 行分片路径。
1D 路径(TRsqrt_1D_Switch)按VFImplKind再细分:
VFIMPL_1D_NO_POST_UPDATE:vlds/vsts使用绝对偏移i * nRepeatElem,不更新基址;VFIMPL_2D_POST_UPDATE/VFIMPL_2D_NO_POST_UPDATE:映射到 2D 内核,按行 stride 寻址;- 默认:
TRsqrt_1D_PostUpdate,vlds/vsts带POST_UPDATE标志,每次 repeat 后自动累加基址。
2D 路径(TRsqrt_2D)对外层行循环、内层列分块循环:每次加载src + i * SrcRowStride + j * nRepeatElem处的数据块,计算后存储到dst + i * DstRowStride + j * nRepeatElem,行与行之间通过编译期常量RowStride跳转,天然适配行主序 tile 的 strided 布局。VFImplKind枚举及后续向量指令下发机制可在 include/pto/npu/a5/vf/vf_defs.hpp 与 include/pto/npu/a5/vf/vf_common.hpp 中进一步追溯。
CPU 侧参考实现(include/pto/cpu/TRSqrt.hpp)则把每个元素提升为double计算1.0 / std::sqrt(x)后再窄化回原类型,并通过cpu::parallel_for_rows按行并行,是跨平台正确性验证(如 tests/cpu/st/testcase/trsqrt)的黄金参照。
测试验证:A5 ST 用例与精度阈值
仓库为 TRSQRT 提供了完整的 ST(System Test)验证,A5 平台用例位于 tests/npu/a5/src/st/testcase/trsqrt,包含三个文件:
- trsqrt_kernel.cpp:内核定义,通过
Event<Op::TLOAD, Op::TRSQRT>与Event<Op::TRSQRT, Op::TSTORE_VEC>构建TLOAD → TRSQRT → TSTORE事件链,并以模板参数highPrecision、isInPlace控制精度与原地/非原地模式; - main.cpp:gtest 驱动,读取
input.bin上板执行后与golden.bin比对; - gen_data.py:生成输入与黄金数据。
测试覆盖了 8 个代表性用例(tests/npu/a5/src/st/testcase/trsqrt/main.cpp):
| 用例 | 类型 | 形态(dst/src) | valid | 模式 |
|---|---|---|---|---|
| case1 | float | 64x64 / 64x64 | 64x64 | 高精度 + 原地 |
| case2 | float | 64x64 / 64x64 | 64x64 | 高精度 + 非原地 |
| case3 | half | 64x64 / 64x64 | 64x64 | 高精度 + 原地 |
| case4 | half | 64x64 / 64x64 | 64x64 | 高精度 + 非原地 |
| case5 | float | 128x128 / 64x64 | 64x64 | 默认 + dst 大于 src |
| case6 | float | 64x64 / 128x128 | 32x32 | 默认 + src 大于 dst |
| case7 | half | 128x256 / 64x64 | 64x64 | 默认 + dst 大于 src |
| case8 | half | 64x64 / 128x256 | 32x32 | 默认 + src 大于 dst |
case5~case8 特意让 dst 与 src 的物理 tile 尺寸不同、valid 区域小于物理尺寸,用以验证有效边界与 stride 寻址的正确性。精度容差(tests/npu/a5/src/st/testcase/trsqrt/main.cpp)按类型与精度分级:float默认eps = 0.00005f,half默认eps = 0.0005f,高精度模式下收紧到eps = 0.0000001f——这从测试侧印证了 HIGH_PRECISION 路径(IEEE 754 除法 + 高精度开方)的精度优势。
除 A5 外,TRSQRT 在 A2A3(tests/npu/a2a3/src/st/testcase/trsqrt)、kirin9030(tests/npu/kirin9030/src/st/testcase/trsqrt)、kirinDev0000(tests/npu/kirinDev0000/src/st/testcase/trsqrt)以及 CPU(tests/cpu/st/testcase/trsqrt)平台均有对应的内核与测试用例,体现了该指令跨平台(cross-platform)的一致性设计;代价模型侧,include/pto/costmodel/a5/vf_costmodel.hpp 也将PtoOpcode::TRSQRT映射到"TRSQRT"名称,供性能仿真与开销估算使用(参见 docs/costmodel/perf-sim-user-guide.md)。
总结
TRSQRT 是 PTO 指令集中实现1/sqrt(x)的基础逐元素向量指令,文档规范、C++ Intrinsic、多平台内核与测试用例共同构成了完整的技术闭环。理解它的关键在于三点:一是其两级汇编语法(SSA/DPS)与 Auto/Manual 两种开发模式的对应关系;二是tmp参数在不同平台(A5 预留兼容 vs A2A3 高精度实际使用)的行为差异;三是 DEFAULT 与 HIGH_PRECISION 两条精度路径(vsqrt+vdiv组合 vs IEEE 754 高精度实现)在精度与开销上的权衡。当你需要编写涉及归一化、数值缩放等场景的 PTO 算子时,可直接参考本文示例与 TRSQRT 指令规范,并借助仓库测试用例完成端到端验证。
【免费下载链接】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),仅供参考