BN3DTrainingReduceGrad 是 3D BatchNorm 训练反向传播的 elementwise 收尾算子:接收损失对 BN 前向输出 y 的梯度 grads、BN 前向输入 x、前置归约段 BN3DTrainingUpdateGrad 产出的逐通道梯度 diff_scale / diff_offset,以及前向统计量 scale / batch_mean / batch_variance,逐元素合成损失对前向输入 x 的梯度 dx(输出 y)。
grads
x
BN3DTrainingUpdateGrad
diff_scale
diff_offset
scale
batch_mean
batch_variance
dx
y
当前 ops-nn 的 norm 目录下已有 1D/2D 归一化训练反向链(如 BNTrainingUpdateGrad / BatchNormGrad 系列),缺少 3D BatchNorm 对应的反向收尾算子,导致 3D 归一化训练反向链路不完整(前向 BN3DTrainingReduce/Update + 反向 BN3DTrainingUpdateGrad 已具备,收尾缺位),需新增本算子补齐。
cann/ops-nn 仓库 norm 目录下新增算子需求
补齐 3D BatchNorm 训练反向链路的收尾段,与前驱 BN3DTrainingUpdateGrad 构成完整的 3D 归一化训练反向算子对(前向 BN3DTrainingReduce → BN3DTrainingUpdate,反向 BN3DTrainingUpdateGrad → BN3DTrainingReduceGrad),支撑 3D 卷积等视频/体积网络训练场景;对齐现有 2D BatchNorm 训练反向收尾的语义与用法,降低上层框架接入成本。
BN3DTrainingReduce
BN3DTrainingUpdate
BN3DTrainingReduceGrad
计算公式:记通道数 C = diff_scale.shape[0],除通道轴外元素个数 num = N·D·H·W,逐通道 s_c = sqrt(batch_variance_c + epsilon):
C = diff_scale.shape[0]
num = N·D·H·W
s_c = sqrt(batch_variance_c + epsilon)
yn,c,d,h,w=(gradsn,c,d,h,w−diff_scalec⋅(xn,c,d,h,w−batch_meanc)num⋅sc−diff_offsetcnum)⋅scalecscy_{n,c,d,h,w} = \left( grads_{n,c,d,h,w} - \frac{diff\_scale_c \cdot (x_{n,c,d,h,w} - batch\_mean_c)}{num \cdot s_c} - \frac{diff\_offset_c}{num} \right) \cdot \frac{scale_c}{s_c} yn,c,d,h,w=(gradsn,c,d,h,w−num⋅scdiff_scalec⋅(xn,c,d,h,w−batch_meanc)−numdiff_offsetc)⋅scscalec
代码逻辑(arch35 / Ascend950,5 条 VF 链逐步对齐 golden):
s = sqrt(bv + eps)
t_a = (x − batch_mean) · diff_scale · inv_num
·inv_num
÷num
t1 = grads − t_a / s
÷s
grads−
t2 = t1 − diff_offset · inv_num
y = (t2 · scale) / s
Backgroud(背景信息)
BN3DTrainingReduceGrad 是 3D BatchNorm 训练反向传播的 elementwise 收尾算子:接收损失对 BN 前向输出 y 的梯度
grads、BN 前向输入x、前置归约段BN3DTrainingUpdateGrad产出的逐通道梯度diff_scale/diff_offset,以及前向统计量scale/batch_mean/batch_variance,逐元素合成损失对前向输入x的梯度dx(输出y)。当前 ops-nn 的 norm 目录下已有 1D/2D 归一化训练反向链(如 BNTrainingUpdateGrad / BatchNormGrad 系列),缺少 3D BatchNorm 对应的反向收尾算子,导致 3D 归一化训练反向链路不完整(前向 BN3DTrainingReduce/Update + 反向 BN3DTrainingUpdateGrad 已具备,收尾缺位),需新增本算子补齐。
Origin(信息来源)
cann/ops-nn 仓库 norm 目录下新增算子需求
Benefit / Necessity(价值/作用)
补齐 3D BatchNorm 训练反向链路的收尾段,与前驱
BN3DTrainingUpdateGrad构成完整的 3D 归一化训练反向算子对(前向BN3DTrainingReduce→BN3DTrainingUpdate,反向BN3DTrainingUpdateGrad→BN3DTrainingReduceGrad),支撑 3D 卷积等视频/体积网络训练场景;对齐现有 2D BatchNorm 训练反向收尾的语义与用法,降低上层框架接入成本。Design(设计方案)
计算公式:记通道数
C = diff_scale.shape[0],除通道轴外元素个数num = N·D·H·W,逐通道s_c = sqrt(batch_variance_c + epsilon):yn,c,d,h,w=(gradsn,c,d,h,w−num⋅scdiff_scalec⋅(xn,c,d,h,w−batch_meanc)−numdiff_offsetc)⋅scscalec
代码逻辑(arch35 / Ascend950,5 条 VF 链逐步对齐 golden):
s = sqrt(bv + eps)(ChainS:Adds + Sqrt);t_a = (x − batch_mean) · diff_scale · inv_num(ChainA1:Sub + Mul + Muls,·inv_num替代÷num);t1 = grads − t_a / s(ChainA2:Div + Sub,÷s在括号内、先于grads−);t2 = t1 − diff_offset · inv_num(ChainB:Muls + Sub);y = (t2 · scale) / s(ChainC:Mul + Div),结果回落到grads的 dtype。grads/x/y∈ {FLOAT16, FLOAT32, BFLOAT16} 三者一致;5 个参数张量恒 FLOAT32;grads/x/y支持 NCDHW(通道轴 dim1)与 NDHWC(通道轴 dim4)双布局;1D 参数沿通道轴广播到 5D,其余轴为广播轴;