PTO-ISA TAXPY 指令深度解析:Tile 级原位缩放累加(a·x+y)的实现原理与实战
2026/9/19 17:56:30 网站建设 项目流程

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必须先初始化(否则累加基是未定义数据);dstsrc0的有效形状必须一一对应(逐元素映射)。

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:

Tiledtype有效形状TileType说明
dsthalffloat$M \times N$Vec(UB)累加基 + 结果(RMW)
src0halffloat$M \times N$Vec(UB)缩放源,逐元素

dstsrc0的有效行数、有效列数必须完全相同。

支持的 dtype 组合

dstdtypesrc0dtypescalardtype说明
halfhalfhalf同类型路径,直接vaxpy
floatfloatfloat同类型路径,直接vaxpy
floathalfhalf差异路径:src0拓宽为 FP32 后累加

dstsrc0必须 dtype 一致,或dstfloatsrc0half(允许 half→float 的拓宽累加)。dsthalfsrc0float的组合非法,由实现内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检查src0dst的 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:

  1. vlds加载src0dst(NORM 模式);
  2. CreatePredicate<T>(sreg)构造尾部谓词掩码,屏蔽不足一个 repeat 的列;
  3. 执行vaxpy(vreg2, vreg0, scalar, preg)
  4. vsts写回dst

差异类型路径(A5 后端:UNPK_B16 + vcvt)

在 include/pto/npu/a5/TAxpy.hpp 的AxpyInstrDiff中,当dstfloatsrc0half时:

  1. vlds(..., UNPK_B16)以半精度解包模式加载src0
  2. vcvt(reg_src_tmp, vreg0, preg, PART_EVEN)将 half 拓宽为 FP32;
  3. 再以(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计算。

约束一览

约束适用范围原因
dstsrc0必须为TileType::Vec所有目标在 UB(向量流水线)上执行
dstsrc0有效形状相同($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),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询