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

RmsNormGradQuant

产品支持情况

产品 是否支持
Ascend 950PR/Ascend 950DT
Atlas A3 训练系列产品/Atlas A3 推理系列产品 ×
Atlas A2 训练系列产品/Atlas A2 推理系列产品 ×
Atlas 200I/500 A2 推理产品 ×
Atlas 推理系列产品 ×
Atlas 训练系列产品 ×

功能说明

  • 算子功能:RmsNormGrad是用于计算RmsNorm的梯度,即在反向传播过程中计算输入张量的梯度的算子。RmsNormGradQuant算子将RmsNormGrad和Quantize两个算子融合,RmsNormGrad计算完dx后进行quant计算,减少搬入搬出操作。

  • 计算公式:

    dxi=(dyi∗gi−1Rms⁡(x)∗xi∗Mean⁡(y))∗1Rms⁡(x), where Mean⁡(y)=1n∑i=1n(dyi∗gi∗xi∗1Rms⁡(x))dx_i= (dy_i * g_i - \frac{1}{\operatorname{Rms}(\mathbf{x})} * x_i * \operatorname{Mean}(\mathbf{y})) * \frac{1} {\operatorname{Rms}(\mathbf{x})}, \quad \text { where } \operatorname{Mean}(\mathbf{y}) = \frac{1}{n}\sum_{i=1}^n (dy_i * g_i * x_i * \frac{1}{\operatorname{Rms}(\mathbf{x})})

    • div_mode为True时:

      dxi_quant=round((dxi/scales_x)+offset_x)dx_i\_quant=round((dx_i / scales\_x) + offset\_x)

    • div_mode为False时:

      dxi_quant=round((dxi∗scales_x)+offset_x)dx_i\_quant=round((dx_i * scales\_x) + offset\_x)

    dgi=1Rms⁡(x)∗xi∗dyidg_i = \frac{1}{\operatorname{Rms}(\mathbf{x})} * x_i * dy_i

参数说明

参数名 输入/输出/属性 描述 数据类型 数据格式
dy 输入
  • 表示反向传回的梯度,对应公式中的dy
  • shape支持2-8维。
FLOAT32、FLOAT16、BFLOAT16 ND
x 输入
  • 表示正向算子的输入,被标准化的数据,对应公式中的x
  • shape和dtype与dy保持一致。
FLOAT32、FLOAT16、BFLOAT16 ND
rstd 输入
  • 表示正向算子的中间计算结果,对应公式中Rms(x)的倒数。
  • shape需要满足rstd_shape = x_shape[0:n],n < x_shape.dims(),n与gamma的n一致。
FLOAT32 ND
gamma 输入
  • 表示正向算子进行归一化计算的缩放因子(权重),对应公式中的g
  • shape需要满足gamma_shape = x_shape[n:],n < x_shape.dims()。
  • dtype与dy相同或为FLOAT32。
FLOAT32、FLOAT16、BFLOAT16 ND
scales_x 输入
  • 表示输入梯度量化缩放因子,对应公式中的scales_x
  • shape为[1],维度为1。
  • dtype与dy相同或为FLOAT32。
FLOAT32、FLOAT16、BFLOAT16 ND
offset_x 可选输入
  • 表示输入梯度量化零点,对应公式中的offset_x
  • shape为[1],维度为1。
INT32 ND
quant_mode 属性
  • 量化模式。
  • 仅支持"static",表示静态量化模式。
STRING -
div_mode 属性
  • 公式中决定量化公式是否使用除法的参数,对应公式中的div_mode
  • 支持True和False。
BOOL -
dst_type 可选属性
  • 表示指定数据转换后dx的类型。
  • 输入范围为{2, 34},分别对应{INT8, HIFLOAT8}。
  • 默认值为2。
INT64 -
dx 输出
  • 表示输入x的量化梯度,对应公式中的dx_quant
  • shape与入参dy的shape保持一致。
INT8、HIFLOAT8 ND
dgamma 输出
  • 表示gamma的梯度,对应公式中的dg
  • shape与入参gamma的shape保持一致。
FLOAT32 ND

约束说明

  • Ascend 950PR/Ascend 950DT:
    • 各输入Tensor支持空Tensor。
    • dy、x、rstd、gamma支持非连续Tensor;scales_x、offset_x不支持非连续Tensor。

调用说明

调用方式 样例代码 说明
aclnn接口 test_aclnn_rms_norm_grad_quant 通过aclnnRmsNormGradQuant接口方式调用RmsNormGradQuant算子。
图模式 - 通过算子IR构图方式调用RmsNormGradQuant算子。