| 文件 | 最后提交记录 | 最后更新时间 |
|---|---|---|
| 6 天前 | ||
| 15 天前 | ||
| 15 天前 | ||
| 15 天前 | ||
| 6 天前 | ||
| 15 天前 | ||
| 3 个月前 | ||
| 1 个月前 |
AddRmsNormDynamicMxQuant
产品支持情况
| 产品 | 是否支持 |
|---|---|
| Ascend 950PR/Ascend 950DT | √ |
| Atlas A3 训练系列产品/Atlas A3 推理系列产品 | × |
| Atlas A2 训练系列产品/Atlas A2 推理系列产品 | × |
| Atlas 200I/500 A2 推理产品 | × |
| Atlas 推理系列产品 | × |
| Atlas 训练系列产品 | × |
| Kirin X90 处理器系列产品 | × |
| Kirin 9030 处理器系列产品 | × |
功能说明
-
算子功能:RmsNorm算子是大模型常用的归一化操作,相比LayerNorm算子,其去掉了减去均值的部分。DynamicMxQuant算子则是在尾轴上按blocksize=32分组进行动态MX量化的算子。AddRmsNormDynamicMxQuant算子将RmsNorm前的Add算子和RmsNorm归一化输出给到的DynamicMxQuant算子融合起来,减少搬入搬出操作。在输入尾轴axis上,根据每blocksize=32个数,计算出这组数对应的量化尺度mxscale,然后对这组数每一个除以mxscale,根据round_mode转换到对应的dst_type,得到量化结果y。在dst_type为FLOAT8_E4M3FN、FLOAT8_E5M2时,根据scale_alg的取值来指定计算mxscale的不同算法。
-
计算公式:
x=x1+x2x=x_{1}+x_{2}
y=RmsNorm(x)=xRms(x)⋅gamma+beta, where Rms(x)=1n∑i=1nxi2+epsilony = \operatorname{RmsNorm}(x)=\frac{x}{\operatorname{Rms}(\mathbf{x})}\cdot gamma+beta, \quad \text { where } \operatorname{Rms}(\mathbf{x})=\sqrt{\frac{1}{n} \sum_{i=1}^n x_i^2+epsilon}
当scale_alg为0时:
- 将RmsNorm输出y在尾轴维度上按k = 32个数分组,一组k个数 {{Vi}i=1k}\{\{V_i\}_{i=1}^{k}\} 动态量化为 {mxscale,{Pi}i=1k}\{mxscale,\{P_i\}_{i=1}^{k}\}
shared_exp=floor(log2(maxi(∣Vi∣)))−emaxshared\_exp = floor(log_2(max_i(|V_i|))) - emax
mxscale=2shared_expmxscale = 2^{shared\_exp}
Pi=cast_to_dst_type(Vi/mxscale,round_mode), i from 1 to blocksizeP_i = cast\_to\_dst\_type(V_i/mxscale, round\_mode), \space i\space from\space 1\space to\space blocksize\\
-
emax: 对应数据类型的最大正则数的指数位。
DataType emax FLOAT4_E2M1 2 FLOAT4_E1M2 0 FLOAT8_E4M3FN 8 FLOAT8_E5M2 15
当scale_alg为1时,只涉及FP8类型:
- 将长向量按块分,每块长度为k,对每块单独计算一个块缩放因子Sfp32bS_{fp32}^b,再把块内所有元素用同一个Sfp32bS_{fp32}^b映射到目标低精度类型FP8。
- 找到该块中数值的最大绝对值:
Amax(Dfp32b)=max({∣di∣}i=1k)Amax(D_{fp32}^b)=max(\{|d_{i}|\}_{i=1}^{k})
- 将FP32映射到目标数据类型FP8可表示的范围内:
Sfp32b=Amax(Dfp32b)Amax(DType)S_{fp32}^b = \frac{Amax(D_{fp32}^b)}{Amax(DType)}
- 转换为FP8格式下可表示的缩放值Sue8m0bS_{ue8m0}^b
- 从块的浮点缩放因子Sfp32bS_{fp32}^b中提取无偏指数EintbE_{int}^b和尾数MfixpbM_{fixp}^b
- 为保证量化时不溢出,对指数进行向上取整:
Eintb={Eintb+1,如果Sfp32b为正规数,且Eintb<254且Mfixpb>0Eintb+1,如果Sfp32b为非正规数,且Mfixpb>0.5Eintb,否则E_{int}^b = \begin{cases} E_{int}^b + 1, & \text{如果} S_{fp32}^b \text{为正规数,且} E_{int}^b < 254 \text{且} M_{fixp}^b > 0 \\ E_{int}^b + 1, & \text{如果} S_{fp32}^b \text{为非正规数,且} M_{fixp}^b > 0.5 \\ E_{int}^b, & \text{否则} \end{cases}
- 计算块缩放因子:Sue8m0b=2EintbS_{ue8m0}^b=2^{E_{int}^b}
- 计算块转换因子:Rfp32b=1fp32(Sue8m0b)R_{fp32}^b=\frac{1}{fp32(S_{ue8m0}^b)}
- 应用到量化的最终步骤:di=DType(dfp32i⋅Rfp32n)d^i = DType(d_{fp32}^i \cdot R_{fp32}^n)
参数说明
| 参数名 | 输入/输出/属性 | 描述 | 数据类型 | 数据格式 |
|---|---|---|---|---|
| x1 | 输入 | 表示标准化过程中的源数据张量,对应公式中的x1。 | FLOAT16、BFLOAT16 | ND |
| x2 | 输入 | 表示标准化过程中的源数据张量,对应公式中的x2。shape和数据类型需要与x1一致。 | FLOAT16、BFLOAT16 | ND |
| gamma | 输入 | 表示标准化过程中的权重张量,对应公式中的gamma。shape需要与x1最后一维一致。 | FLOAT16、BFLOAT16、FLOAT32 | ND |
| beta | 可选输入 | 表示标准化过程中的偏置项,对应公式中的beta。shape必须与gamma一致。 | FLOAT16、BFLOAT16、FLOAT32 | ND |
| epsilon | 可选属性 |
|
FLOAT32 | - |
| scale_alg | 可选属性 |
|
INT64 | - |
| round_mode | 可选属性 |
|
STRING | - |
| dst_type | 可选属性 |
|
INT64 | - |
| output_rstd | 可选属性 |
|
BOOL | - |
| y | 输出 |
|
FLOAT4_E2M1、FLOAT4_E1M2、FLOAT8_E4M3FN、FLOAT8_E5M2 | ND |
| x | 输出 |
|
FLOAT16、BFLOAT16 | ND |
| mxscale | 输出 |
|
FLOAT8_E8M0 | ND |
| rstd | 输出 |
|
FLOAT32 | ND |
约束说明
-
Ascend 950PR/Ascend 950DT:
mxscale的shape约束说明如下:
- rank(mxscale) = rank(x1) + 1。
- mxscale.shape[-2] = (ceil(x1.shape[-1] / 32) + 2 - 1) / 2。
- mxscale.shape[-1] = 2。
- 其他维度与输入x1一致。
-
当输出y的数据类型为FLOAT4_E2M1或FLOAT4_E1M2,x1尾轴的值必须为偶数。
-
输入gamma、可选输入beta的数据类型只能和x1的数据类型保持一致或者为FLOAT32。
-
边界值场景说明
- 当输入是Inf时:1、输出y为0;2、输出x为Inf;3、输出mxscale为255,偶数pad填充值为0;4、输出rstd为0。
- 当输入是NaN时:1、输出y为0;2、输出x为Nan;3、输出mxscale为255,偶数pad填充值为0;4、输出rstd为NaN。
调用说明
| 调用方式 | 样例代码 | 说明 |
|---|---|---|
| aclnn接口 | test_aclnn_add_rms_norm_dynamic_mx_quant | 通过aclnnAddRmsNormDynamicMxQuant接口方式调用AddRmsNormDynamicMxQuant算子。 |
| 图模式 | - | 通过算子IR构图方式调用AddRmsNormDynamicMxQuant算子。 |