PTO-ISA TAXPY 指令深度解析:Tile 级原位缩放累加(a·x+y)的实现原理与实战
【免费下载链接】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 为核心,完整讲解 PTO(Parallel 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.md):
dst:既是累加输入($y$)也是输出,调用前必须已初始化;src0:只读($x$),逐元素参与运算;scalar:标量缩放系数($a$),类型为TileDataSrc::DType。
这与普通的TADDS(Tile 加标量)不同:TAXPY是两个 Tile 之间带标量系数的融合运算,输入输出共用一个 Tile(RMW,读-修改-写),天然适合在已有累加基上做带权累加。
数学语义:基于有效区域的逐元素 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.hpp:
template <typename TileDataDst, typename TileDataSrc, typename... WaitEvents> PTO_INST RecordEvent TAXPY(TileDataDst &dst, TileDataSrc &src0, typename TileDataSrc::DType scalar, WaitEvents &...events);| 参数 | 方向 | 含义 |
|---|---|---|
dst | 输入/输出 | 累加基与结果 Tile($y$),读-修改-写,Vec |
src0 | 输入 | 缩放源 Tile($x$),只读,Vec,有效形状与dst相同 |
scalar | 输入 | 标量缩放系数($a$),类型为TileDataSrc::DType |
events... | 输入 | 等待事件(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$ 的 Tile:
| Tile | dtype | 有效形状 | TileType | 说明 |
|---|---|---|---|---|
dst | half或float | $M \times N$ | Vec(UB) | 累加基 + 结果(RMW) |
src0 | half或float | $M \times N$ | Vec(UB) | 缩放源,逐元素 |
dst与src0的有效行数、有效列数必须完全相同。
支持的 dtype 组合
dstdtype | src0dtype | scalardtype | 说明 |
|---|---|---|---|
half | half | half | 同类型路径,直接vaxpy |
float | float | float | 同类型路径,直接vaxpy |
float | half | half | 差异路径: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_assert:
static_assert(std::is_same_v<T, half> || std::is_same_v<T, float>, "TAXPY: Invalid data type"); static_assert(std::is_same_v<T, U> || (std::is_same_v<T, float> && std::is_same_v<U, 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 一致。
底层实现原理:向量流水线上的 vaxpy
TAXPY在向量流水线(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)),然后逐行、逐 repeat:
vlds加载src0与dst(NORM 模式);- 用
CreatePredicate<T>(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)以半精度解包模式加载src0;vcvt(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上限 255);validCol / elementsPerRepeat > validRow判定列 repeat 数与行数的关系。
随后二选一:
- count 模式(
AxpyCountMode,include/pto/npu/a2a3/TAxpy.hpp):set_mask_count()+SetVectorCount(validCol),逐行执行一次vaxpy(dstPtr, src0Ptr, scalar, 0, 1, 1, 8, 4)——注意差异类型下 src 一个 repeat 只占 4 个 block,dst 占 8 个 block,由vaxpy原生处理 half→float 的 4-block src / 8-block dst 映射; - norm 模式(
AxpyNormMode/AxpyNormModeTail,include/pto/npu/a2a3/TAxpy.hpp):先按REPEAT_MAX切分行块,行内按dstElementsPerRepeat循环切列,剩余列用SetContMaskByDType<U>(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_cast<typename TileDataDst::DType>(src.data()[srcIdx] * scalar); volatile auto result = static_cast<typename 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 = Shape<1, 1, 1, vRows, vCols>; using DynStridDim5 = pto::Stride<1, 1, 1, vCols, 1>; using GlobalData = GlobalTensor<T, DynShapeDim5, DynStridDim5>; using TileData = Tile<TileType::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/(A5)
- tests/npu/a2a3/src/st/testcase/taxpy/(A2/A3)
- tests/npu/kirin9030/src/st/testcase/taxpy/(Kirin9030)
- tests/cpu/st/testcase/taxpy/(CPU 参考实现)
相关资源
- 指令文档:docs/isa/TAXPY.md / docs/isa/TAXPY_zh.md
- 公共声明与包装函数:include/pto/common/pto_instr.hpp
- A5 后端实现:include/pto/npu/a5/TAxpy.hpp
- A2/A3 后端实现:include/pto/npu/a2a3/TAxpy.hpp
- CPU_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),仅供参考