| 文件 | 最后提交记录 | 最后更新时间 |
|---|---|---|
| 1 个月前 | ||
| 5 天前 | ||
| 13 天前 | ||
| 22 天前 | ||
| 12 天前 | ||
| 22 天前 | ||
| 1 个月前 | ||
| 1 个月前 |
InstanceNormGrad
产品支持情况
| 产品 | 是否支持 |
|---|---|
| Ascend 950PR/Ascend 950DT | √ |
| Atlas A3 训练系列产品/Atlas A3 推理系列产品 | √ |
| Atlas A2 训练系列产品/Atlas A2 推理系列产品 | √ |
| Atlas 200I/500 A2 推理产品 | √ |
| Atlas 推理系列产品 | √ |
| Atlas 训练系列产品 | √ |
功能说明
-
算子功能:Instance Normalization的反向传播。给定上游梯度
dy与前向的x、variance、mean、gamma,计算pd_x、pd_gamma、pd_beta。pd_x逐(N, C)在空间维(D, H, W)上归约;pd_gamma、pd_beta在空间维之外再对N归约,只保留C维。 -
计算公式(
variance为原始方差,ε为编译期常量1e-6,rstd = (variance + ε)^(-1/2),m = D*H*W,R = {D,H,W},x_hat = (x - mean) * rstd):pd_xl=dy⋅gammapd\_xl = dy \cdot gamma
pd_var=∑R(−0.5⋅pd_xl⋅(x−mean)⋅(variance+ε)−3/2)pd\_var = \sum_{R} \left( -0.5 \cdot pd\_xl \cdot (x - mean) \cdot (variance + \varepsilon)^{-3/2} \right)
pd_mean=∑R(−1.0⋅pd_xl⋅rstd)pd\_mean = \sum_{R} \left( -1.0 \cdot pd\_xl \cdot rstd \right)
pd_x=pd_xl⋅rstd+pd_var⋅2m⋅(x−mean)+pd_mean⋅1mpd\_x = pd\_xl \cdot rstd + pd\_var \cdot \frac{2}{m} \cdot (x - mean) + pd\_mean \cdot \frac{1}{m}
pd_gamma=∑N,R(dy⋅x_hat),pd_beta=∑N,R(dy)pd\_gamma = \sum_{N,R} (dy \cdot x\_hat), \quad pd\_beta = \sum_{N,R} (dy)
参数说明
| 参数名 | 输入/输出/属性 | 描述 | 数据类型 | 数据格式 |
|---|---|---|---|---|
| dy | 输入 | 表示反向传回的梯度,对应公式中的`dy`。shape与数据类型与入参`x`一致。 | FLOAT16、FLOAT32 | NDHWC |
| x | 输入 | 表示正向算子的输入,对应公式中的`x`。 | FLOAT16、FLOAT32 | NDHWC |
| variance | 输入 | 表示每个instance的方差,对应公式中的`variance`。数据类型与入参`dy`一致。 | FLOAT16、FLOAT32 | NDHWC |
| mean | 输入 | 表示每个instance的均值,对应公式中的`mean`。shape与数据类型与入参`variance`一致。 | FLOAT16、FLOAT32 | NDHWC |
| gamma | 输入 | 表示标准化过程中的缩放张量,对应公式中的`gamma`。数据类型与入参`dy`一致。 | FLOAT16、FLOAT32 | NDHWC |
| pd_x | 输出 | 表示对`x`的梯度,对应公式中的`pd_x`。shape与数据类型与入参`x`一致。 | FLOAT16、FLOAT32 | NDHWC |
| pd_gamma | 输出 | 表示对`gamma`的梯度,对应公式中的`pd_gamma`。shape与数据类型与入参`gamma`一致。 | FLOAT16、FLOAT32 | NDHWC |
| pd_beta | 输出 | 表示对`beta`的梯度,对应公式中的`pd_beta`。shape与数据类型与入参`gamma`一致。 | FLOAT16、FLOAT32 | NDHWC |
约束说明
- 确定性说明:确定性实现。
x为空Tensor时,pd_x同为空Tensor,pd_gamma、pd_beta的所有元素均为0。
调用说明
| 调用方式 | 样例代码 | 说明 |
|---|---|---|
| 图模式 | test_geir_instance_norm_grad | 通过算子IR构图方式调用InstanceNormGrad算子。 |