文件最后提交记录最后更新时间
6 天前
15 天前
15 天前
15 天前
6 天前
15 天前
3 个月前
1 个月前
README

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 可选属性
  • 表示添加到分母中的值,以确保数值稳定。对应公式中的epsilon。
  • 默认值为1e-6。
FLOAT32 -
scale_alg 可选属性
  • 表示mxscale的计算方法,对应公式中的scale_alg。
  • 支持取值0和1,取值为0表示Open Compute Project(OCP)实现,取值为1表示cuBLAS实现。当dst_type为FLOAT4_E2M1/FLOAT4_E1M2时仅支持取值为0。
  • 默认值为0。
INT64 -
round_mode 可选属性
  • 表示数据转换的模式,对应公式中的round_mode。
  • 当dst_type为40/41时,支持{"rint", "floor", "round"}。
  • 当dst_type为36/35时,仅支持{"rint"}。
  • 默认值为"rint"。
STRING -
dst_type 可选属性
  • 表示指定数据转换后y的类型,对应公式中的DType。
  • 输入范围为{35, 36, 40, 41},分别对应{FLOAT8_E5M2, FLOAT8_E4M3FN, FLOAT4_E2M1, FLOAT4_E1M2}。
  • 默认值为40。
INT64 -
output_rstd 可选属性
  • 表示指定是否输出有效的rstd_out。
  • 支持True和False。
  • 默认值为False。
  • 当output_rstd为False时,rstd为无效占位输出。
BOOL -
y 输出
  • 表示归一化并量化后的结果,对应公式中的Pi和di,shape与x1一致。
FLOAT4_E2M1、FLOAT4_E1M2、FLOAT8_E4M3FN、FLOAT8_E5M2 ND
x 输出
  • 表示x1和x2的和,对应公式中的x。shape和数据类型与x1一致。
FLOAT16、BFLOAT16 ND
mxscale 输出
  • 表示每个分组对应的量化尺度,对应公式中的mxscale和Sb,shape见约束说明。
FLOAT8_E8M0 ND
rstd 输出
  • 表示归一化后的标准差的倒数,对应公式中Rms(x)的倒数。
  • 当output_rstd为True时,shape与入参x1的shape前几维保持一致,前几维指x1的维度减去gamma的维度,表示不需要norm的维度。
  • 当output_rstd为False时,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算子。