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

GroupNormSwishGrad

产品支持情况

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

功能说明

dYTempi=xi^⋅gamma+betadYTemp_i = \hat{x_i} \cdot gamma + beta

dSwishTempi=swishScale⋅dYTempi−swishScale⋅dYTempiexp⁡(−swishScale⋅dYTempi)+1+1dSwishTemp_i = swishScale \cdot dYTemp_i - \frac{swishScale \cdot dYTemp_i}{\exp(-swishScale \cdot dYTemp_i) + 1} + 1

dYNewi=dSwishTempiexp⁡(−swishScale⋅dYTempi)+1⋅dydYNew_i = \frac{dSwishTemp_i}{\exp(-swishScale \cdot dYTemp_i) + 1} \cdot dy

dBeta=∑i=1ndYNewidBeta = \sum_{i=1}^n dYNew_i

dGamma=∑i=1n(dYNewi⋅xi^)dGamma = \sum_{i=1}^n (dYNew_i \cdot \hat{x_i})

dx=rstd⋅(dYNew∗gamma−x^∗(∑i=1ngammai∗dGamma)−(∑i=1ngammai∗dBeta))dx = rstd \cdot (dYNew * gamma - \hat{x} * (\sum_{i=1}^n gamma_i * dGamma) - (\sum_{i=1}^n gamma_i * dBeta))

参数说明

参数名 输入/输出/属性 描述 数据类型 数据格式
dy 输入 反向计算的梯度。 FLOAT32、FLOAT16、BFLOAT16 ND
mean 输入 正向计算的第二个输出,表示input分组后每个组的均值。 FLOAT32、FLOAT16、BFLOAT16 ND
rstd 输入 正向计算的第三个输出,表示input分组后每个组的标准差倒数。 FLOAT32、FLOAT16、BFLOAT16 ND
x 输入 正向的输入x。 FLOAT32、FLOAT16、BFLOAT16 ND
gamma 输入 每个channel的缩放系数。 FLOAT32、FLOAT16、BFLOAT16 ND
beta 输入 表示每个channel的偏移系数。 FLOAT32、FLOAT16、BFLOAT16 ND
numGroups 属性 表示将输入gradOut的C维度分为group组。 INT64 -
dataFormatOptional 属性 数据格式。 Char* -
swishScale 属性 Swish计算公式中的系数。 Float -
dgammaIsRequire 属性 是否需要输出dgamma。 Bool -
dbetaIsRequire 属性 是否需要输出dbeta。 Bool -
dxOut 输出 x的输出梯度。 FLOAT32、FLOAT16、BFLOAT16 -
dgammaOut 输出 gamma的输出梯度。 FLOAT32、FLOAT16、BFLOAT16 -
dbetaOut 输出 beta的输出梯度。 FLOAT32、FLOAT16、BFLOAT16 -

约束说明

调用说明

调用方式 样例代码 说明
aclnn接口 test_aclnn_group_norm_swish_grad.cpp 通过aclnnGroupNormSwishGrad接口方式调用GroupNormSwishGrad算子。
图模式 - 通过算子IR构图方式调用GroupNormSwishGrad算子。