Skip to content

MmadBitMode

产品支持情况

不传入bias的原型

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

传入bias的原型

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

功能说明

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

MmadBitMode对于MmadParams结构体构造进行了优化,本接口适用于scalar流水成为性能优化瓶颈的场景,支持基础Mmad/MmadMx计算功能。本接口与Mmad/MmadMx接口的差异在于参数传入的方式不同,本接口传入的是联合结构体MmadBitModeParams。MmadBitModeParams类参数设计思想说明:

联合体(union)是一种特殊的数据结构,允许在相同的内存位置存储不同的数据类型。union的所有成员共享同一块内存空间,大小由最大成员决定,同一时间只能使用一个成员。

位域(bit-field)是一种特殊的类成员,允许精确控制结构体中成员变量所占用的内存位数。结构体中成员变量从上到下对应内存中从低位到高位。

MmadBitModeParams类使用union与bit-field方法,采用bit位表达参数类型,使用bit-field结构体自动处理入参的bit位数,并利用union的特性实现多参数融合传递,仅需传递一个入参即可包含全部所需信息,对应底层接口仅需要接收一个参数。同时,当需要修改参数中某一bit位的值时,仅需要通过循环和位运算即可实现,不需要重新传入参数,减少了scalar计算,实现性能提升。

MmadBitModeParams类可以直接使用MmadBitModeParams结构体类型对象初始化:

C++
__aicore__ inline MmadBitModeParams(const MmadBitModeParams &mmadParams_);

也可以使用各参数的Set函数修改参数值,并且由于使用了联合体,还可以对config0直接进行逐bit位修改来修改参数。

函数原型

  • 不传入bias

    C++
    template <typename T, typename U, typename S>
    __aicore__ inline void Mmad(const LocalTensor<T>& dst, const LocalTensor<U>& fm, const LocalTensor<S>& filter, const MmadBitModeParams& mmadParams)
    
    template <typename T, typename U, typename S>
    __aicore__ inline void MmadMx(const LocalTensor<T>& dst, const LocalTensor<U>& fm, const LocalTensor<S>& filter, const MmadBitModeParams& mmadParams)
    
  • 传入bias

    C++
    template <typename T, typename U, typename S, typename V>
    __aicore__ inline void Mmad(const LocalTensor<T>& dst, const LocalTensor<U>& fm, const LocalTensor<S>& filter, const LocalTensor<V>& bias, const MmadBitModeParams& mmadParams)
    
    template <typename T, typename U, typename S, typename V>
    __aicore__ inline void MmadMx(const LocalTensor<T>& dst, const LocalTensor<U>& fm, const LocalTensor<S>& filter, const LocalTensor<V>& bias, const MmadBitModeParams& mmadParams)
    

参数说明

表1 参数说明

参数名称输入/输出含义
dst输出目的操作数,结果矩阵C,类型为LocalTensor,支持的物理存储位置为L0C Buffer(TPosition:CO1)。
LocalTensor的起始地址需要按照1024字节对齐。
fm输入源操作数,左矩阵A,类型为LocalTensor,支持的物理存储位置为L0A Buffer(TPosition:A2)。
左矩阵A对应的scale矩阵起始地址为:A矩阵起始对应地址/16。
对于fp4场景LocalTensor的起始地址需要按照512字节对齐。对于fp8场景LocalTensor的起始地址需要按照1024字节对齐。
filter输入源操作数,右矩阵B,类型为LocalTensor,支持的物理存储位置为L0B Buffer(TPosition:B2)。
右矩阵b对应的scale矩阵起始地址为:B矩阵起始对应地址/16。
对于fp4场景LocalTensor的起始地址需要按照512字节对齐。对于fp8场景LocalTensor的起始地址需要按照1024字节对齐。
bias输入源操作数,Bias矩阵,类型为LocalTensor,支持的物理存储位置为BT Buffer(TPosition:C2)。
LocalTensor的起始地址需要按照64字节对齐。
mmadParams输入矩阵乘相关参数。
该参数类型的具体定义请参考${INSTALL_DIR}/asc/include/basic_api/kernel_struct_mm.h${INSTALL_DIR}请替换为CANN软件安装后文件存储路径。
MmadBitModeParams参数说明请参考下表。

表2 MmadBitModeParams类参数说明

参数名称含义
config0uint64_t类型,与MmadBitModeConfig0位域(bit-field)结构体类型参数config0BitMode组成联合体(union),初始化为0,可以使用类对象的GetConfig0()函数获取其值。
config0BitModeMmadBitModeConfig0位域(bit-field)结构体类型,参数参考表3,与config0组成联合体(union)。

表3 MmadBitModeConfig0结构体参数说明

参数名称含义
m左矩阵Height,取值范围:m∈[0, 4095]。默认值为0。
该参数是位域结构体的最低位参数,占用12bit,可以使用MmadBitModeParams类对象的SetM()函数设置其值,使用GetM()函数获取其值。
k左矩阵Width、右矩阵Height,取值范围:k∈[0, 4095]。默认值为0。
该参数是位域结构体的第二低位参数,占用12bit,可以使用MmadBitModeParams类对象的SetK()函数设置其值,使用GetK()函数获取其值。
n右矩阵Width,取值范围:n∈[0, 4095]。默认值为0。
该参数是位域结构体的第三低位参数,占用12bit,可以使用MmadBitModeParams类对象的SetN()函数设置其值,使用GetN()函数获取其值。
unitFlag预留参数。为后续的功能做保留,开发者暂时无需关注,使用默认值即可。
该参数是位域结构体的第四低位参数,占用2bit,可以使用MmadBitModeParams类对象的SetUnitFlag()函数设置其值,使用GetUnitFlag()函数获取其值。
disableGemvM = 1时,用于配置Mmad计算是否开启GEMV。当输入为false时,表示开启GEMV;反之,输入为true时,表示关闭GEMV。
GEMV(General Matrix-Vector Multiplication)表示实现矩阵和向量的乘积,开启GEMV后,Mmad API从L0A Buffer读取数据时,数据将以ND格式进行读取,而不会将其视为ZZ格式。
该参数是位域结构体的第五低位参数,占用1bit,可以使用MmadBitModeParams类对象的SetDisableGemv()函数设置其值,使用GetDisableGemv()函数获取其值。
cmatrixSource配置C矩阵初始值是否来源于BT Buffer(TPosition:C2)。默认值为false。
  • false:来源于L0C Buffer(TPosition:CO1);
  • true:来源于BT Buffer(TPosition:C2)。
注意:带bias输入的接口配置该参数无效,会根据bias输入的位置来判断C矩阵初始值是否来源于L0C Buffer还是BT Buffer。
该参数是位域结构体的第六低位参数,占用1bit,可以使用MmadBitModeParams类对象的SetCmatrixSource()函数设置其值,使用GetCmatrixSource()函数获取其值。
cmatrixInitVal配置C矩阵初始值是否为0。默认值为 true。
  • true:C矩阵初始值为0;
  • false:C矩阵初始值通过cmatrixSource参数进行配置。
该参数是位域结构体的最高位参数,占用1bit,可以使用MmadBitModeParams类对象的SetCmatrixInitVal()函数设置其值,使用GetCmatrixInitVal()函数获取其值。

数据类型

表4 Mmad接口左矩阵、右矩阵、Bias矩阵、结果矩阵支持的精度类型组合

左矩阵fm type右矩阵filter typebias type结果矩阵dst type
int8_tint8_tint32_tint32_t
halfhalffloatfloat
floatfloatfloatfloat
bfloat16_tbfloat16_tfloatfloat
fp8_e4m3fn_tfp8_e4m3fn_tfloatfloat
fp8_e4m3fn_tfp8_e5m2_tfloatfloat
fp8_e5m2_tfp8_e4m3fn_tfloatfloat
fp8_e5m2_tfp8_e5m2_tfloatfloat
hifloat8_thifloat8_tfloatfloat

表5 MmadMx接口左矩阵、右矩阵、Scale矩阵、Bias矩阵、结果矩阵支持的精度类型组合

左矩阵fm右矩阵filterScale矩阵偏置Bias结果矩阵dst
fp4x2_e1m2_tfp4x2_e1m2_tfp8_e8m0_tfloatfloat
fp4x2_e2m1_tfp4x2_e1m2_tfp8_e8m0_tfloatfloat
fp4x2_e1m2_tfp4x2_e2m1_tfp8_e8m0_tfloatfloat
fp4x2_e2m1_tfp4x2_e2m1_tfp8_e8m0_tfloatfloat
fp8_e4m3fn_tfp8_e4m3fn_tfp8_e8m0_tfloatfloat
fp8_e4m3fn_tfp8_e5m2_tfp8_e8m0_tfloatfloat
fp8_e5m2_tfp8_e4m3fn_tfp8_e8m0_tfloatfloat
fp8_e5m2_tfp8_e5m2_tfp8_e8m0_tfloatfloat

返回值说明

约束说明

  • 不同矩阵对于存储位置的约束:
    • 结果矩阵C只支持位于物理存储位置为L0C Buffer(TPosition:CO1),大小256KB。
    • 左矩阵A只支持位于物理存储位置为L0A Buffer(TPosition:A2),大小64KB。
    • 右矩阵B只支持位于物理存储位置为L0B Buffer(TPosition:B2),大小64KB。
    • Bias矩阵只支持位于物理存储位置为BT Buffer(TPosition:C2),大小4KB。
  • 地址约束说明请参考表1。
  • 当M、K、N中的任意一个值为0时,表示指令不会执行,该接口将被视为NOP(空操作)。
  • K需要是64的倍数。
  • 对于fp4场景A/B矩阵的起始地址需要按照512字节对齐,对于fp8场景A/B矩阵的起始地址需要按照1024字节对齐。
  • 左矩阵A/B对应的scale矩阵起始地址为:A/B矩阵起始对应地址/16。
  • 其他特殊场景约束可参考Mmad接口约束说明

调用示例

样例请参考Mmad样例

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