TSCATTER¶
指令示意图¶
简介¶
TSCATTER提供两种操作模式:
- 索引散播(Index-based Scatter):使用逐元素行索引将源Tile的行散播到目标Tile中。
- 掩码散播(Mask Scatter):按照掩码模式将源元素散播到目标位置,并在元素间交错填充零值。支持按行散播(
SCATTER_ROW)和按列散播(SCATTER_COL)两种模式。
数学语义¶
索引散播¶
对每个源元素 (i, j),写入:
\[ \mathrm{dst}_{\mathrm{idx}_{i,j},\ j} = \mathrm{src}_{i,j} \]
若多个元素映射到同一目标位置,最终值由实现定义(当前实现中以最后写入者为准)。
掩码散播¶
对于掩码模式 P,将源元素散播并交错填充零值。散播方向由 ScatterAxis 控制:
SCATTER_ROW(默认)¶
沿列方向散播,扩展列维度:
\[ \mathrm{dst}_{i, P \cdot j + \mathrm{pos}_P} = \mathrm{src}_{i,j} \]
\[ \mathrm{dst}_{i, P \cdot j + \mathrm{zeros}_P} = 0 \]
其中:
DstTileData::ValidCol=SrcTileData::ValidCol× 扩展倍数DstTileData::ValidRow=SrcTileData::ValidRow
SCATTER_COL¶
沿行方向散播,扩展行维度:
\[ \mathrm{dst}_{P \cdot i + \mathrm{pos}_P, j} = \mathrm{src}_{i,j} \]
\[ \mathrm{dst}_{P \cdot i + \mathrm{zeros}_P, j} = 0 \]
其中:
DstTileData::ValidRow=SrcTileData::ValidRow× 扩展倍数DstTileData::ValidCol=SrcTileData::ValidCol
扩展倍数¶
P1010或P0101:扩展倍数 = 2P0001、P0010、P0100、P1000:扩展倍数 = 4P1111:扩展倍数 = 1(等同于TMOV)
汇编语法¶
同步形式:
%dst = tscatter %src, %idx : !pto.tile<...>, !pto.tile<...> -> !pto.tile<...>
AS Level 1(SSA)¶
%dst = pto.tscatter %src, %idx : (!pto.tile<...>, !pto.tile<...>) -> !pto.tile<...>
AS Level 2(DPS)¶
pto.tscatter ins(%src, %idx : !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 TileDataD, typename TileDataS, typename TileDataI, typename... WaitEvents>
PTO_INST RecordEvent TSCATTER(TileDataD& dst, TileDataS& src, TileDataI& indexes, WaitEvents&... events);
掩码散播¶
template <MaskPattern maskPattern = MaskPattern::P1111, auto ScatterType = ScatterAxis::SCATTER_ROW,
typename DstTileData, typename SrcTileData, typename... WaitEvents>
PTO_INST RecordEvent TSCATTER(DstTileData& dst, SrcTileData& src, WaitEvents&... events);
MaskPattern枚举¶
定义于 include/pto/common/type.hpp:
| 值 | 模式 | 描述 | 扩展倍数 |
|---|---|---|---|
P0101 |
01010101... | 每两个元素取第一个 | ×2 |
P1010 |
10101010... | 每两个元素取第二个 | ×2 |
P0001 |
00010001... | 每四个元素取第一个 | ×4 |
P0010 |
00100010... | 每四个元素取第二个 | ×4 |
P0100 |
01000100... | 每四个元素取第三个 | ×4 |
P1000 |
10001000... | 每四个元素取第四个 | ×4 |
P1111 |
11111111... | 取全部元素(等同于TMOV) | ×1 |
ScatterAxis枚举¶
定义于 include/pto/common/type.hpp:
| 值 | 描述 |
|---|---|
SCATTER_ROW |
沿列方向散播,扩展列维度(默认) |
SCATTER_COL |
沿行方向散播,扩展行维度 |
约束¶
索引散播¶
- 实现检查 (Atlas A2/A3 训练系列产品/Atlas A2/A3 推理系列产品):
TileDataD::Loc、TileDataS::Loc、TileDataI::Loc必须是TileType::Vec。TileDataD::DType、TileDataS::DType必须是以下之一:int32_t、int16_t、int8_t、half、float16_t、float32_t、uint32_t、uint16_t、uint8_t、bfloat16_t。TileDataI::DType必须是以下之一:int16_t、int32_t、uint16_t或uint32_t。- 不对
indexes值执行边界检查。 - 静态有效边界:
TileDataD::ValidRow <= TileDataD::Rows、TileDataD::ValidCol <= TileDataD::Cols、TileDataS::ValidRow <= TileDataS::Rows、TileDataS::ValidCol <= TileDataS::Cols、TileDataI::ValidRow <= TileDataI::Rows、TileDataI::ValidCol <= TileDataI::Cols。 TileDataD::DType与TileDataS::DType必须相同。- 当
TileDataD::DType大小为4字节时,TileDataI::DType大小必须为4字节。 - 当
TileDataD::DType大小为2字节时,TileDataI::DType大小必须为2字节。 - 当
TileDataD::DType大小为1字节时,TileDataI::DType大小必须为2字节。
- 实现检查 (Ascend 950PR/Ascend 950DT):
TileDataD::Loc、TileDataS::Loc、TileDataI::Loc必须是TileType::Vec。TileDataD::DType、TileDataS::DType必须是以下之一:int32_t、int16_t、int8_t、half、float16_t、float32_t、uint32_t、uint16_t、uint8_t、bfloat16_t。TileDataI::DType必须是以下之一:int16_t、int32_t、uint16_t或uint32_t。- 不对
indexes值执行边界检查。 - 静态有效边界:
TileDataD::ValidRow <= TileDataD::Rows、TileDataD::ValidCol <= TileDataD::Cols、TileDataS::ValidRow <= TileDataS::Rows、TileDataS::ValidCol <= TileDataS::Cols、TileDataI::ValidRow <= TileDataI::Rows、TileDataI::ValidCol <= TileDataI::Cols。 TileDataD::DType与TileDataS::DType必须相同。- 当
TileDataD::DType大小为4字节时,TileDataI::DType大小必须为4字节。 - 当
TileDataD::DType大小为2字节时,TileDataI::DType大小必须为2字节。 - 当
TileDataD::DType大小为1字节时,TileDataI::DType大小必须为2字节。
掩码散播¶
- 实现检查 (Atlas A2/A3 训练系列产品/Atlas A2/A3 推理系列产品):
DstTileData::Loc、SrcTileData::Loc必须是TileType::Vec。DstTileData::DType、SrcTileData::DType必须是以下之一:int32_t、int16_t、int8_t、half、float16_t、float32_t、uint32_t、uint16_t、uint8_t、bfloat16_t。DstTileData::DType与SrcTileData::DType必须相同。maskPattern必须在P0101到P1111范围内。- 静态有效边界:
DstTileData::ValidCol <= DstTileData::Cols、SrcTileData::ValidCol <= SrcTileData::Cols、DstTileData::ValidRow <= DstTileData::Rows、SrcTileData::ValidRow <= SrcTileData::Rows。 P1111模式等同TMOV:要求validRow和validCol分别匹配,内部通过TMOV_IMPL实现。
- 实现检查 (Ascend 950PR/Ascend 950DT):
DstTileData::Loc、SrcTileData::Loc必须是TileType::Vec。DstTileData::DType、SrcTileData::DType必须是以下之一:int32_t、int16_t、int8_t、half、float16_t、float32_t、uint32_t、uint16_t、uint8_t、bfloat16_t。DstTileData::DType与SrcTileData::DType必须相同。maskPattern必须在P0101到P1111范围内。- 静态有效边界:
DstTileData::ValidRow <= DstTileData::Rows、DstTileData::ValidCol <= DstTileData::Cols、SrcTileData::ValidRow <= SrcTileData::Rows、SrcTileData::ValidCol <= SrcTileData::Cols。 SCATTER_ROW模式运行时断言:SrcTileData::ValidRow必须等于DstTileData::ValidRow。SrcTileData::ValidCol必须等于DstTileData::ValidCol × 扩展倍数,扩展倍数取决于掩码模式(P1111为1,P1010/P0101为2,P0001/P0010/P0100/P1000为4)。
SCATTER_COL模式运行时断言:SrcTileData::ValidCol必须等于DstTileData::ValidCol。SrcTileData::ValidRow必须等于DstTileData::ValidRow × 扩展倍数,扩展倍数取决于掩码模式(P1111为1,P1010/P0101为2,P0001/P0010/P0100/P1000为4)。
重要提示¶
警告:在执行散播操作前,目标Tile缓冲区会完全初始化为0(整个Tile大小
Rows × Cols),不受ValidRow和ValidCol限制。这意味着:
- 分配给
dstTile的整个UB缓冲区都会被写入零值。ValidRow/ValidCol范围之外的元素在操作后也将为零。- 请确保目标Tile的UB缓冲区不会与其他活跃数据重叠。
示例¶
索引散播(自动Auto)¶
#include <pto/pto-inst.hpp>
using namespace pto;
void example_auto() {
using TileT = Tile<TileType::Vec, float, 16, 16>;
using IdxT = Tile<TileType::Vec, uint16_t, 16, 16>;
TileT src, dst;
IdxT idx;
TSCATTER(dst, src, idx);
}
索引散播(手动Manual)¶
#include <pto/pto-inst.hpp>
using namespace pto;
void example_manual() {
using TileT = Tile<TileType::Vec, float, 16, 16>;
using IdxT = Tile<TileType::Vec, uint16_t, 16, 16>;
TileT src, dst;
IdxT idx;
TASSIGN(src, 0x1000);
TASSIGN(dst, 0x2000);
TASSIGN(idx, 0x3000);
TSCATTER(dst, src, idx);
}
掩码散播(自动Auto)¶
#include <pto/pto-inst.hpp>
using namespace pto;
void example_mask_auto() {
// P1010: 目标大小 = 源大小 × 2
using SrcTileT = Tile<TileType::Vec, half, 16, 64>;
using DstTileT = Tile<TileType::Vec, half, 16, 128>;
SrcTileT src;
DstTileT dst;
TSCATTER<MaskPattern::P1010>(dst, src);
}
void example_mask_p1000() {
// P1000: 目标大小 = 源大小 × 4
using SrcTileT = Tile<TileType::Vec, float, 16, 64>;
using DstTileT = Tile<TileType::Vec, float, 16, 256>;
SrcTileT src;
DstTileT dst;
TSCATTER<MaskPattern::P1000>(dst, src);
}
void example_mask_scatter_col() {
// SCATTER_COL: 沿行方向散播,扩展行维度
// P1010: 目标行数 = 源行数 × 2
using SrcTileT = Tile<TileType::Vec, half, 64, 16>;
using DstTileT = Tile<TileType::Vec, half, 128, 16>;
SrcTileT src;
DstTileT dst;
TSCATTER<MaskPattern::P1010, ScatterAxis::SCATTER_COL>(dst, src);
}
掩码散播(手动Manual)¶
#include <pto/pto-inst.hpp>
using namespace pto;
void example_mask_manual() {
using SrcTileT = Tile<TileType::Vec, half, 16, 64>;
using DstTileT = Tile<TileType::Vec, half, 16, 128>;
SrcTileT src;
DstTileT dst;
TASSIGN(src, 0x1000);
TASSIGN(dst, 0x2000);
TSCATTER<MaskPattern::P1010>(dst, src);
}
void example_mask_manual_scatter_col() {
// SCATTER_COL 手动绑定模式
using SrcTileT = Tile<TileType::Vec, half, 64, 16>;
using DstTileT = Tile<TileType::Vec, half, 128, 16>;
SrcTileT src;
DstTileT dst;
TASSIGN(src, 0x1000);
TASSIGN(dst, 0x2000);
TSCATTER<MaskPattern::P1010, ScatterAxis::SCATTER_COL>(dst, src);
}
汇编示例(ASM)¶
自动模式¶
# 自动模式:由编译器/运行时负责资源放置与调度。
%dst = pto.tscatter %src, %idx : (!pto.tile<...>, !pto.tile<...>) -> !pto.tile<...>
手动模式¶
# 手动模式:先显式绑定资源,再发射指令。
# 可选(当该指令包含 tile 操作数时):
# pto.tassign %arg0, @tile(0x1000)
# pto.tassign %arg1, @tile(0x2000)
%dst = pto.tscatter %src, %idx : (!pto.tile<...>, !pto.tile<...>) -> !pto.tile<...>
PTO汇编形式¶
%dst = tscatter %src, %idx : !pto.tile<...>, !pto.tile<...> -> !pto.tile<...>
# AS Level 2 (DPS)
pto.tscatter ins(%src, %idx : !pto.tile_buf<...>, !pto.tile_buf<...>) outs(%dst : !pto.tile_buf<...>)