| 文件 | 最后提交记录 | 最后更新时间 |
|---|---|---|
| 4 小时前 | ||
| 15 天前 | ||
| 6 个月前 | ||
| 15 天前 | ||
| 15 天前 | ||
| 15 天前 | ||
| 7 个月前 | ||
| 1 个月前 |
GroupNormSwishGrad
产品支持情况
| 产品 | 是否支持 |
|---|---|
| Ascend 950PR/Ascend 950DT | √ |
| Atlas A3 训练系列产品/Atlas A3 推理系列产品 | √ |
| Atlas A2 训练系列产品/Atlas A2 推理系列产品 | √ |
| Atlas 200I/500 A2 推理产品 | × |
| Atlas 推理系列产品 | × |
| Atlas 训练系列产品 | × |
功能说明
-
算子功能:aclnnGroupNormSwish的反向操作。
-
计算公式:
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算子。 |