文件最后提交记录最后更新时间
20 天前
1 个月前
7 个月前
20 天前
24 天前
1 个月前
7 个月前
1 个月前
README

GroupNormGrad

产品支持情况

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

功能说明

  • 算子功能:GroupNorm用于计算输入张量的组归一化结果,均值,标准差的倒数,该算子是对GroupNorm的反向计算。用于计算输入张量的梯度,以便在反向传播过程中更新模型参数。

  • 计算公式:

    x^=(x−mean)⋅rstd\hat{x} = (x - mean) \cdot rstd

    dβ=∑i=1ndyd\beta = \sum_{i=1}^n dy

    dγ=∑i=1n(dy⋅x^)d\gamma = \sum_{i=1}^n (dy \cdot \hat{x})

    dx=rstd⋅γ[dy−1N(dβ+x^⋅dγ)]dx = rstd \cdot \gamma \begin{bmatrix} dy - \frac{1}{N} (d\beta + \hat{x} \cdot d\gamma) \end{bmatrix}

参数说明

参数名 输入/输出/属性 描述 数据类型 数据格式
dy 输入 反向计算的梯度tensor,对应公式中的`dy`。数据类型与`x`相同。`dy`支持2-8维(N, C, *),计算逻辑仅关注前两个维度(N和C),其余维度可合并为一个维度。 FLOAT16、FLOAT32、BFLOAT16 ND
mean 输入 正向计算的第二个输出,表示`x`分组后每个组的均值,对应公式中的`mean`。必须是2D(N, num_groups)。 FLOAT16、FLOAT32、BFLOAT16 ND
rstd 输入 正向计算的第三个输出,表示`x`分组后每个组的标准差倒数,对应公式中的`rstd`。数据类型与`mean`相同。必须是2D(N, num_groups)。 FLOAT16、FLOAT32、BFLOAT16 ND
x 输入 正向计算的首个输入,对应公式中的`x`。数据类型与`dy`相同。支持2-8维(N, C, *),计算逻辑仅关注前两个维度(N和C),其余维度可合并为一个维度。 FLOAT16、FLOAT32、BFLOAT16 ND
gamma 输入 表示每个channel的缩放系数,对应公式中的`γ`,数据类型与`mean`相同。必须是1D。`gamma`的值需要与`x`的C轴值一致。 FLOAT16、FLOAT32、BFLOAT16 ND
num_groups 属性 表示将输入`dy`的C维度分为group组,group需大于0。 INT -
data_format 可选属性
  • 指定输出的`dx`的数据格式。
  • 默认值为NCHW。
STRING -
dx_is_require 可选属性
  • 是否输出`dx`。
  • 默认值为true。
BOOL -
dgamma_is_require 可选属性
  • 是否输出`dgamma`。
  • 默认值为true。
BOOL -
dbeta_is_require 可选属性
  • 是否输出`dbeta`。
  • 默认值为true。
BOOL -
dx 输出 计算输出的梯度,用于更新输入`x`的梯度,对应公式中的`dx`。数据类型与`dy`相同,shape与`x`相同。 FLOAT16、FLOAT32、BFLOAT16 ND
dgamma 输出 计算输出的梯度,用于更新缩放参数的梯度,对应公式中的`dγ`。数据类型与`mean`相同,shape与`gamma`相同。 FLOAT16、FLOAT32、BFLOAT16 ND
dbeta 输出 计算输出的梯度,用于更新偏置参数的梯度,对应公式中的`dβ`。数据类型与`mean`相同,shape与`gamma`相同。 FLOAT16、FLOAT32、BFLOAT16 ND
  • Atlas 训练系列产品:输入和输出参数的数据类型不支持BFLOAT16。

约束说明

调用说明

调用方式 样例代码 说明
aclnn接口 test_aclnn_group_norm_grad 通过aclnnGroupNormBackward接口方式调用GroupNormGrad算子。
图模式 - 通过算子IR构图方式调用GroupNormGrad算子。