文件最后提交记录最后更新时间
1 个月前
5 天前
13 天前
22 天前
12 天前
22 天前
1 个月前
1 个月前
README

InstanceNormGrad

产品支持情况

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

功能说明

  • 算子功能:Instance Normalization的反向传播。给定上游梯度dy与前向的xvariancemeangamma,计算pd_xpd_gammapd_betapd_x(N, C)在空间维(D, H, W)上归约;pd_gammapd_beta在空间维之外再对N归约,只保留C维。

  • 计算公式(variance为原始方差,ε为编译期常量1e-6rstd = (variance + ε)^(-1/2)m = D*H*WR = {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_gammapd_beta的所有元素均为0。

调用说明

调用方式 样例代码 说明
图模式 test_geir_instance_norm_grad 通过算子IR构图方式调用InstanceNormGrad算子。