TAXPY¶
简介¶
对Tile执行原位缩放累加(AXPY,\(a \cdot x + y\)):将 src0 按标量 scalar 缩放后累加到 dst 上。
\[ \mathrm{dst}_{i,j} \leftarrow \mathrm{scalar} \cdot \mathrm{src0}_{i,j} + \mathrm{dst}_{i,j} \]
dst 既是累加输入(\(y\))也是输出,调用前必须已初始化;src0(\(x\))只读;scalar(\(a\))为标量。
数学语义¶
对于有效区域中的每个元素 (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。
除非另有说明,语义在有效区域内定义,目标相关行为标记为实现定义。
C++内建接口¶
声明于 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),指令前隐式 TSYNC |
Tile尺寸与数据类型¶
对于有效形状 \(M \times N\):
| Tile | dtype | 有效形状 | TileType | 说明 |
|---|---|---|---|---|
dst |
half 或 float |
\(M \times N\) | Vec (UB) |
累加基 + 结果(RMW) |
src0 |
half 或 float |
\(M \times N\) | Vec (UB) |
缩放源,逐元素 |
dst与src0的有效行数、有效列数必须完全相同。
支持的输入dtype¶
dst dtype |
src0 dtype |
scalar dtype |
说明 |
|---|---|---|---|
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拦截)。
实现说明¶
TAXPY在向量流水线(PIPE_V)上执行,使用 vaxpy(\(a \cdot x + y\))向量内建:
- 同类型(
dst与src0同dtype):逐repeat加载src0与dst,执行vaxpy(dst, src0, scalar)后写回dst;尾部不足一个repeat的列由谓词掩码屏蔽。 - 差异类型(
dst=float,src0=half):src0的half数据拓宽为FP32后参与累加(Ascend 950PR/Ascend 950DT上经UNPK_B16解包并vcvt转换;Atlas A2/A3 训练系列产品/Atlas A2/A3 推理系列产品由vaxpy原生按4-block src / 8-block dst处理)。 - Atlas A2/A3 训练系列产品/Atlas A2/A3 推理系列产品上按repeat-stride是否溢出、以及列数与行数的关系,在count模式与norm模式间选择,以覆盖任意有效形状。
约束¶
| 约束 | 适用范围 | 原因 |
|---|---|---|
dst、src0 必须为 TileType::Vec |
所有目标 | 在UB(向量流水线)上执行 |
dst 与 src0 有效形状相同(\(M \times N\)) |
所有目标 | 逐元素一一对应 |
dst dtype ∈ {half, float} |
所有目标 | vaxpy 支持的浮点字宽 |
dst/src0 dtype一致,或 (float,half) |
所有目标 | 仅允许half→float拓宽累加 |
dst 调用前必须已初始化 |
所有目标 | dst 作为累加基 \(y\) 被读入 |
示例¶
// dst 必须先初始化(作为累加基 y);结果:dst = scalar * src0 + dst
TAXPY(dstTile, srcTile, scalar);
完整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参考实现)。