接口功能:融合GroupedMatmul、activation和quant,详细解释见计算公式。此接口为WeightNz特化版本,调用者必须传入FRACTAL_NZ格式的weight,接口按该格式解析该参数。当前版本仅支持MX量化场景,激活函数仅支持gelu_tanh。
计算公式:
[object Object]定义:
⋅ 表示矩阵乘法。
⊙ 表示逐元素乘法。
gelu_tanh激活函数的数学语义为:
当前kernel底层实现使用gelu_sigmoid函数对gelu_tanh进行近似计算,定义为:
* 表示专家数, 表示总token数, 表示输入特征维度, 表示输出特征维度。
- 表示MX量化时共享指数的分组大小,当前仅支持64。
输入:
- :激活矩阵。
- :分组索引列表,按groupListType解释为cumsum或count。
- :分组weight矩阵。
- :weight矩阵量化因子。
- :MXFP8场景下必须为空,支持nullptr、空tensorlist或长度为1且元素shape为(0)的空tensorlist。
- :左矩阵量化因子,对应公式中的。
输出:
- :激活并量化后的输出矩阵,数据类型由输出Tensor y指定。
- :输出量化因子。
计算过程
- 根据groupList[i]确定当前分组的token范围,。
- 根据分组确定的入参进行GroupedMatmul和反量化计算,中间GroupedMatmul结果默认为FLOAT32类型:
- 执行gelu_tanh激活:
注:当前kernel底层实现使用gelu_sigmoid函数近似计算gelu_tanh,近似公式见定义。
- 对激活结果进行MX量化,目标数据类型DType由输出Tensor y的数据类型指定:
场景1,当scaleAlg为0时,表示OCP实现,将激活结果在N轴按分组,一组个数动态量化为。
量化后的按对应的位置组成输出,按对应N轴分组组成输出。
:对应数据类型的最大正则数的指数位。
[object Object]undefined
场景2,当scaleAlg为1时,表示cuBLAS实现,只涉及FP8类型。将激活结果在N轴按分组,每块单独计算一个块缩放因子,再把块内所有元素用同一个映射到目标FP8类型。如果最后一块不足个元素,缺失值视为0并按完整块处理。
找到该块中数值的最大绝对值:
将FP32映射到目标数据类型FP8可表示的范围内,其中是目标精度能表示的最大值:
将块缩放因子转换为FP8格式下可表示的缩放值。先从中提取无偏指数和尾数,再为避免量化溢出对指数向上取整:
对块内每个元素执行量化:
最终输出的量化结果为,其中表示块缩放因子,表示块内量化后的数据。
每个算子分为,必须先调用“aclnnGroupedMatmulActivationQuantWeightNzGetWorkspaceSize”接口获取计算所需workspace大小以及包含了算子计算流程的执行器,再调用“aclnnGroupedMatmulActivationQuantWeightNz”接口执行计算。
参数说明
[object Object]- [object Object]Ascend 950PR/Ascend 950DT[object Object]:
- x仅支持非转置输入,weight支持非转置和转置输入,接口会根据weight和weightScale的形状及转置关系推导weight是否转置。
- weightScale转置属性需要与weight保持一致。
- 空Tensor处理规则:
- weight和weightScale作为必选输入,不支持空tensorlist,tensorlist中的元素不能为nullptr。
- 支持M为0或N为0的空Tensor场景:x和xScaleOptional支持M为0,weight和weightScale支持N为0;该场景下允许K为0,第一段接口返回ACLNN_SUCCESS,workspaceSize为0;当M和N均不为0时,不支持K为0。
- 支持groupList为空的空Tensor场景:当groupList为空,且输出y或yScale为空时,第一段接口返回ACLNN_SUCCESS,workspaceSize为0。
- [object Object]Ascend 950PR/Ascend 950DT[object Object]:
返回值
第一段接口完成入参校验,出现以下场景时报错:
[object Object]
- 确定性计算:
- aclnnGroupedMatmulActivationQuantWeightNz默认确定性实现。
- 非空Tensor场景下,groupList第1维最小为1,最大为1024。
- MXFP8量化场景下需满足以下约束条件:
数据类型需要满足下表:
[object Object]shape约束需要满足下表:
[object Object]N必须为64整数倍。