Skip to content

离散搬入(Gather)

产品支持情况

  • Ascend 950PR/Ascend 950DT:支持
  • Atlas A3 训练系列产品/Atlas A3 推理系列产品:不支持
  • Atlas A2 训练系列产品/Atlas A2 推理系列产品:不支持
  • Atlas 200I/500 A2 推理产品:不支持
  • Atlas 推理系列产品AI Core:不支持
  • Atlas 推理系列产品Vector Core:不支持
  • Atlas 训练系列产品:不支持

功能说明

头文件路径为:"basic_api/reg_compute/kernel_reg_compute_datacopy_intf.h"

该指令会根据索引值index将源操作数按元素收集到目的操作数dstReg中。收集过程如图1所示:

图 1 Gather功能说明

图1 Gather功能说明

函数原型

C++
template <typename T0 = DefaultType, typename T1, typename T2 = DefaultType, typename T3, typename T4>
__simd_callee__ inline void Gather(T3& dstReg, __ubuf__ T1* baseAddr, T4& index, MaskReg& mask)

参数说明

表 1 模板参数说明

参数名描述
T0目的操作数的数据类型。支持的数据类型请参考数据类型
T1源操作数的数据类型。支持的数据类型请参考数据类型
T2索引值的数据类型,支持的数据类型请参考数据类型
T3目的操作数的RegTensor类型。例如RegTensor<half>,由编译器自动推导,用户不需要手动填写。
T4索引值的RegTensor类型,例如RegTensor<uint16_t>,由编译器自动推导,用户不需要手动填写。

表 2 参数说明

参数名输入/输出描述
dstReg输出目的操作数,类型为RegTensor
baseAddr输入源操作数,UB中的基地址,需要32字节对齐。
index输入索引值,dstReg中的每个元素在UB中相对于baseAddr的位置,单位:元素。类型为RegTensor。index中的值可以重复。
例:baseAddr:[elem0, elem1, elem2, elem3, elem4, elem5, elem6, elem7, ...]。
每个元素相对于baseAddr的索引位置为:[0, 1, 2, 3, 4, 5, 6, 7, ...]。
mask输入源操作数元素操作的有效指示,详细说明请参考MaskReg

数据类型

表 3 Gather操作数数据类型对应表

目的操作数源操作数索引值
int16_tint8_tuint16_t
int16_tint16_tuint16_t
uint16_tuint8_tuint16_t
uint16_tuint16_tuint16_t
halfhalfuint16_t
bfloat16_tbfloat16_tuint16_t
int32_tint32_tuint32_t
uint32_tuint32_tuint32_t
floatfloatuint32_t
int64_tint64_tuint32_t
int64_tint64_tuint64_t
uint64_tuint64_tuint32_t
uint64_tuint64_tuint64_t

返回值说明

约束说明

  • 位于Unified Buffer的地址必须32字节对齐。

  • index索引值对应的数据必须在UB有效地址范围内。

  • 当T0为b16类型,T1为b8数据类型时,目的操作数的低8位与源操作数相同,高8位自动补0。例如源操作数T1数据类型为int8_t:

    40=0b00101000 -> 0b0000000000101000,扩充至16位后等于40;

    -40=0b11011000 -> 0b0000000011011000,扩充至16位后等于216。

  • 当源操作数数据类型T1为B64时,T0、T1、T2、T3、T4只支持以下组合:

    T0T1T2T3(自动推导)T4(自动推导)备注
    int64_tint64_tuint32_tRegTensor<int64_t, RegTraitNumOne>RegTensor<uint32_t>index的前32个数有效
    int64_tint64_tuint32_tRegTensor<int64_t, RegTraitNumTwo>RegTensor<uint32_t>-
    int64_tint64_tuint64_tRegTensor<int64_t, RegTraitNumOne>RegTensor<uint64_t, RegTraitNumOne>-
    int64_tint64_tuint64_tRegTensor<int64_t, RegTraitNumOne>RegTensor<uint64_t, RegTraitNumTwo>index的前32个数有效
    int64_tint64_tuint64_tRegTensor<int64_t, RegTraitNumTwo>RegTensor<uint64_t, RegTraitNumTwo>-
    uint64_tuint64_tuint32_tRegTensor<uint64_t, RegTraitNumOne>RegTensor<uint32_t>index的前32个数有效
    uint64_tuint64_tuint32_tRegTensor<uint64_t, RegTraitNumTwo>RegTensor<uint32_t>-
    uint64_tuint64_tuint64_tRegTensor<uint64_t, RegTraitNumOne>RegTensor<uint64_t, RegTraitNumOne>-
    uint64_tuint64_tuint64_tRegTensor<uint64_t, RegTraitNumOne>RegTensor<uint64_t, RegTraitNumTwo>index的前32个数有效
    uint64_tuint64_tuint64_tRegTensor<uint64_t, RegTraitNumTwo>RegTensor<uint64_t, RegTraitNumTwo>-

调用示例

C++
template <typename T, typename U>
__simd_vf__ inline void GatherVF(__ubuf__ T* dstAddr, __ubuf__ T* srcAddr, __ubuf__ U* indexAddr, uint32_t count, uint16_t oneRepeatSize)
{
    AscendC::Reg::RegTensor<T> dstReg;
    AscendC::Reg::RegTensor<U> indexReg;
    AscendC::Reg::MaskReg mask;
    uint16_t repeatTimes = AscendC::CeilDivision(count, oneRepeatSize);
    for (uint16_t i = 0; i < repeatTimes; ++i) {
        mask = AscendC::Reg::UpdateMask<T>(count);
        AscendC::Reg::LoadAlign(indexReg, indexAddr + i * oneRepeatSize);
        AscendC::Reg::Gather(dstReg, srcAddr, indexReg, mask);
        AscendC::Reg::StoreAlign(dstAddr + i * oneRepeatSize, dstReg, mask);
    }
}

免责声明:本站内容由 asc-devkit 仓 master 分支自动编译生成,属于持续开发版本,可能存在缺陷,仅供预览与参考。如需稳定及商用资料,请查阅官方 昇腾社区