TROWARGMAX

指令示意图

TROWARGMAX tile operation

简介

获取每行最大值对应列索引,或同时获取每行最大值及其对应列索引。

数学语义

R = src.GetValidRow()C = src.GetValidCol()。对 0 <= i < R

\[ \mathrm{dst}_{i,0} = \underset{0 \le j < C}{\operatorname{argmax}} \; \mathrm{src}_{i,j} \]
\[ \mathrm{dstval}_{i,0} = \max_{0 \le j < C} \mathrm{src}_{i,j} \]

汇编语法

同步形式:

%dst = trowargmax %src : !pto.tile<...> -> !pto.tile<...>

Lowering may introduce internal scratch tiles; the C++ intrinsic requires an explicit tmp operand.

IR Level 1(SSA)

%dst = pto.trowargmax %src, %tmp : (!pto.tile<...>, !pto.tile<...>) -> !pto.tile<...>

IR Level 2(DPS)

pto.trowargmax ins(%src, %tmp : !pto.tile_buf<...>, !pto.tile_buf<...>) outs(%dst : !pto.tile_buf<...>)

C++内建接口

声明于 include/pto/common/pto_instr.hpp

公共包含头为 <pto/pto-inst.hpp>,内部声明位于 pto/common/pto_instr.hpp

仅输出索引:

template <typename TileDataOut, typename TileDataIn, typename TileDataTmp, typename... WaitEvents>
PTO_INST RecordEvent TROWARGMAX(TileDataOut& dst, TileDataIn& src, TileDataTmp& tmp, WaitEvents&... events);

同时输出值和索引:

template <typename TileDataOutVal, typename TileDataOutIdx, typename TileDataIn, typename TileDataTmp,
          typename... WaitEvents>
PTO_INST RecordEvent TROWARGMAX(TileDataOutVal &dstVal, TileDataOutIdx &dstIdx, TileDataIn &src, TileDataTmp &tmp,
                                WaitEvents &... events)

约束

通用约束或检查

  • 支持的源元素类型:halffloatint32_tint16_t(A2A3)。A5 接受任意 2/4 字节源(halfbfloat16_tint16_tuint16_tfloatint32_tuint32_t)。
  • src 必须使用标准ND布局:行主且非分形(BLayout::RowMajorSLayout::NoneBox)。
  • 仅输出索引时:
    • dstsrc 必须为 TileType::Vec
    • 支持的目标元素类型:uint32_tint32_t
    • 运行时检查遵循共享的行归约检查路径:
      • src.GetValidRow() != 0
      • src.GetValidCol() != 0
      • src.GetValidRow() == dst.GetValidRow()
    • dst 通过共享的行归约索引检查路径约束,可使用以下任一非分形布局:
      • 单列DN布局(BLayout::ColMajorCols == 1),或
      • 有效列数为1的ND布局。
  • 同时输出值和索引时:
    • dstValdstIdxsrc 必须为 TileType::Vec
    • dstVal的元素类型必须与src的元素类型一致。
    • 支持的目标元素类型:
      • 源元素类型为float时,支持uint32_tint32_t
      • 源元素类型为half时,支持uint16_tint16_t
    • 运行时检查遵循共享的行归约检查路径:
      • src.GetValidRow() != 0
      • src.GetValidCol() != 0
      • src.GetValidRow() == dstIdx.GetValidRow()
      • src.GetValidRow() == dstVal.GetValidRow()
    • dstValdstIdx通过共享的行归约索引检查路径约束,可使用以下任一非分形布局:
      • 单列DN布局(BLayout::ColMajorCols == 1),或
      • 有效列数为1的ND布局。

tmp临时Tile相关说明

  • 仅Atlas A2/A3 训练系列产品/Atlas A2/A3 推理系列产品使用tmp临时Tile,Ascend 950PR/Ascend 950DT接收tmp但实际并不使用。
  • Atlas A2/A3 训练系列产品/Atlas A2/A3 推理系列产品实现根据 srcValidColelementPerRepeat(以下缩写为 elemPerRpt)的关系选择三条代码路径之一:

情况1:srcValidCol <= elemPerRpt

  • 仅输出索引模式tmp 不使用。硬件 vcmax 指令直接写入 dst
  • 值+索引模式tmp 用作小型缓冲区(每行2个元素:一个值 + 一个索引)。tmp 可使用以下任一非分形布局:
    • 单列DN布局(BLayout::ColMajorCols == 1),有效行数为 srcValidRow * 2
    • 有效行数为 srcValidRow 且有效列数为2的ND布局。

情况2:elemPerRpt < srcValidCol <= elemPerRpt²(单阶段归约)

  • tmp 被使用于单阶段归约。
  • tmp tile的行数与 src 相同。
  • tmp tile每行所需stride按以下公式计算:
R1 = ceil(validCol / elemPerRpt)
stride = (ceil(R1 * 2 / elemPerBlock) + ceil(R1 / elemPerBlock)) * elemPerBlock

情况3:srcValidCol > elemPerRpt²(两阶段归约)

  • tmp 被使用于两阶段归约,所需空间大于单阶段。
  • tmp tile的行数与 src 相同。
  • tmp tile每行所需stride按以下公式计算:
R1 = ceil(validCol / elemPerRpt)
R2 = ceil(R1 / elemPerRpt)
stage1_size = ceil(R1 * 2 / elemPerBlock) * elemPerBlock
stage2_end  = ceil(R1 / elemPerBlock) * elemPerBlock + ceil(R2 * 2 / elemPerBlock) * elemPerBlock
stride = max(stage1_size, stage2_end) + 2
  • + 2 用于存放每行tmp区域末尾的最终值+索引结果。

示例

自动(Auto)

#include <pto/pto-inst.hpp>

using namespace pto;

void example_auto() {
  using SrcT = Tile<TileType::Vec, float, 16, 16>;
  using DstT = Tile<TileType::Vec, uint32_t, 16, 1, BLayout::ColMajor>;
  using DstValT = Tile<TileType::Vec, float, 16, 1, BLayout::ColMajor>;
  using TmpT = Tile<TileType::Vec, float, 16, 16>;
  SrcT src;
  DstT dst;
  DstValT dstVal;
  TmpT tmp;
  TROWARGMAX(dst, src, tmp);
  TROWARGMAX(dstVal, dst, src, tmp);
}

手动(Manual)

#include <pto/pto-inst.hpp>

using namespace pto;

void example_manual() {
  using SrcT = Tile<TileType::Vec, float, 16, 16>;
  using DstT = Tile<TileType::Vec, uint32_t, 16, 1, BLayout::ColMajor>;
  using DstValT = Tile<TileType::Vec, float, 16, 1, BLayout::ColMajor>;
  using TmpT = Tile<TileType::Vec, float, 16, 16>;
  SrcT src;
  DstT dstIdx;
  DstValT dstVal;
  TmpT tmp;
  TASSIGN(src, 0x1000);
  TASSIGN(dstIdx, 0x2000);
  TASSIGN(dstVal, 0x3000);
  TASSIGN(tmp, 0x4000);
  TROWARGMAX(dstIdx, src, tmp);
  TROWARGMAX(dstVal, dstIdx, src, tmp);
}

汇编示例(ASM)

自动模式

# 自动模式:由编译器/运行时负责资源放置与调度。
%dst = pto.trowargmax %src, %tmp : (!pto.tile<...>, !pto.tile<...>) -> !pto.tile<...>

手动模式

# 手动模式:先显式绑定资源,再发射指令。
# 可选(当该指令包含 tile 操作数时):
# pto.tassign %arg0, @tile(0x1000)
# pto.tassign %arg1, @tile(0x2000)
%dst = pto.trowargmax %src, %tmp : (!pto.tile<...>, !pto.tile<...>) -> !pto.tile<...>

PTO汇编形式

%dst = trowargmax %src : !pto.tile<...> -> !pto.tile<...>
# IR Level 2 (DPS)
pto.trowargmax ins(%src, %tmp : !pto.tile_buf<...>, !pto.tile_buf<...>) outs(%dst : !pto.tile_buf<...>)