Skip to content

Gemm(废弃)

产品支持情况

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

功能说明

该接口废弃,并将在后续版本移除,请不要使用该接口。

根据输入的切分规则,将给定的两个输入张量做矩阵乘,输出至结果张量。将A和B两个输入矩阵乘法在一起,得到一个输出矩阵C。

函数原型

  • 功能接口:

    C++
    template <typename T, typename U, typename S>
    __aicore__ inline void Gemm(const LocalTensor<T>& dst, const LocalTensor<U>& src0, const LocalTensor<S>& src1, const uint32_t m, const uint32_t k, const uint32_t n, GemmTiling tiling, bool partialsum = true, int32_t initValue = 0)
    
  • 切分方案计算接口:

    C++
    template <typename T>
    __aicore__ inline GemmTiling GetGemmTiling(uint32_t m, uint32_t k, uint32_t n)
    

参数说明

表1 接口参数说明

参数名称类型说明
dst输出目的操作数。

Atlas 训练系列产品,支持的TPosition为:CO1,CO2

Atlas 推理系列产品AI Core,支持的TPosition为:CO1,CO2
src0输入源操作数,TPosition为A1。
src1输入源操作数,TPosition为B1。
m输入左矩阵Src0Local有效Height,范围:[1, 4096]。
注意:m可以不是16的倍数。
k输入左矩阵Src0Local有效Width、右矩阵Src1Local有效Height。
•当输入张量Src0Local的数据类型为float时,范围:[1, 8192]
•当输入张量Src0Local的数据类型为half时,范围:[1, 16384]
•当输入张量Src0Local的数据类型为int8_t时,范围:[1, 32768]

注意:k可以不是16的倍数。
n输入右矩阵Src1Local有效Width,范围:[1, 4096]。
注意:n可以不是16的倍数。
tiling输入切分规则,类型为GemmTiling,结构体具体定义为:

struct GemmTiling {
const uint32_t blockSize = 16;
LoopMode loopMode = LoopMode::MODE_NM;
uint32_t mNum = 0;
uint32_t nNum = 0;
uint32_t kNum = 0;
uint32_t roundM = 0;
uint32_t roundN = 0;
uint32_t roundK = 0;
uint32_t c0Size = 32;
uint32_t dtypeSize = 1;
uint32_t mBlockNum = 0;
uint32_t nBlockNum = 0;
uint32_t kBlockNum = 0;
uint32_t mIterNum = 0;
uint32_t nIterNum = 0;
uint32_t kIterNum = 0;
uint32_t mTileBlock = 0;
uint32_t nTileBlock = 0;
uint32_t kTileBlock = 0;
uint32_t kTailBlock = 0;
uint32_t mTailBlock = 0;
uint32_t nTailBlock = 0;
bool kHasTail = false;
bool mHasTail = false;
bool nHasTail = false;
bool kHasTailEle = false;
uint32_t kTailEle = 0;
};

参数说明请参考表3
partialsum输入当dst参数所在的TPosition为CO2时,通过该参数控制计算结果是否搬出。
•取值0:搬出计算结果;
•取值1:不搬出计算结果,可以进行后续计算。
initValue输入表示dst是否需要初始化。
•取值0:dst需要初始化,dst初始矩阵保存有之前结果,新计算结果会累加前一次Gemm计算结果。
•取值1:dst不需要初始化,dst初始矩阵中数据无意义,计算结果直接覆盖dst中的数据。

表2 src0、src1和dst的数据类型组合

src0.dtypesrc1.dtypedst.dtype
int8_tint8_tint32_t
halfhalffloat
halfhalfhalf

表3 GemmTiling结构内参数说明

参数名称类型说明
blockSizeuint32_t固定值,恒为16,一个维度内存放的元素个数。
loopModeLoopMode遍历模式,结构体具体定义为:


enum class LoopMode {
MODE_NM = 0,
MODE_MN = 1,
MODE_KM = 2,
MODE_KN = 3
};
mNumuint32_tM轴等效数据长度参数值,范围:[1, 4096]。
nNumuint32_tN轴等效数据长度参数值,范围:[1, 4096]。
kNumuint32_tK轴等效数据长度参数值。
•当输入张量Src0Local的数据类型为float时,范围:[1, 8192]。
•当输入张量Src0Local的数据类型为half时,范围:[1, 16384]。
•当输入张量Src0Local的数据类型为int8_t时,范围:[1, 32768]。
roundMuint32_tM轴等效数据长度参数值且以blockSize为倍数向上取整,范围:[1, 4096]。
roundNuint32_tN轴等效数据长度参数值且以blockSize为倍数向上取整,范围:[1, 4096]。
roundKuint32_tK轴等效数据长度参数值且以c0Size为倍数向上取整。
•当输入张量Src0Local的数据类型为float时,范围:[1, 8192]。
•当输入张量Src0Local的数据类型为half时,范围:[1, 16384]。
•当输入张量Src0Local的数据类型为int8_t时,范围:[1, 32768]。
c0Sizeuint32_t一个block的字节长度,范围:[16或者32]。
dtypeSizeuint32_t传入的数据类型的字节长度,范围:[1, 2]。
mBlockNumuint32_tM轴Block个数,mBlockNum = mNum / blockSize。
nBlockNumuint32_tN轴Block个数,nBlockNum = nNum / blockSize。
kBlockNumuint32_tK轴Block个数,kBlockNum = kNum / blockSize。
mIterNumuint32_t遍历M轴维度数量,范围:[1, 4096]。
nIterNumuint32_t遍历N轴维度数量,范围:[1, 4096]。
kIterNumuint32_t遍历K轴维度数量,范围:[1, 4096]。
mTileBlockuint32_tM轴切分块个数,范围:[1, 4096]。
nTileBlockuint32_tN轴切分块个数,范围:[1, 4096]。
kTileBlockuint32_tK轴切分块个数,范围:[1, 4096]。
kTailBlockuint32_tK轴尾块个数,范围:[1, 4096]。
mTailBlockuint32_tM轴尾块个数,范围:[1, 4096]。
nTailBlockuint32_tN轴尾块个数,范围:[1, 4096]。
kHasTailboolK轴是否存在尾块。
mHasTailboolM轴是否存在尾块。
nHasTailboolN轴是否存在尾块。
kHasTailElebool是否存在尾块元素。
kTailEleuint32_tK轴尾块元素,范围:[1, 4096]。

数据类型

表4 src0、src1和dst的数据类型组合

src0.dtypesrc1.dtypedst.dtype
int8_tint8_tint32_t
halfhalffloat
halfhalfhalf

返回值说明

约束说明

  • 参数m,k,n可以不是16对齐,但因硬件原因,操作数dst,Src0Local和Src1Local的shape需满足对齐要求,即m方向,n方向要求向上16对齐,k方向根据操作数数据类型按16或32向上对齐。
  • 操作数地址对齐要求请参见通用地址对齐约束

调用示例

该接口已废弃,请使用Mmad接口替代。

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