已关闭
[Requirement|需求建议]: 下一代支持BN3DTrainingReduceGrad #5223
Sun创建于  9 天前关闭于  8 天前
Sun
Sun成员
9 天前 创建

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 归一化训练反向算子对(前向 BN3DTrainingReduceBN3DTrainingUpdate,反向 BN3DTrainingUpdateGradBN3DTrainingReduceGrad),支撑 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,wdiff_scalec(xn,c,d,h,wbatch_meanc)numscdiff_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}

代码逻辑(arch35 / Ascend950,5 条 VF 链逐步对齐 golden):

  1. s = sqrt(bv + eps)(ChainS:Adds + Sqrt);
  2. t_a = (x − batch_mean) · diff_scale · inv_num(ChainA1:Sub + Mul + Muls,·inv_num 替代 ÷num);
  3. t1 = grads − t_a / s(ChainA2:Div + Sub,÷s 在括号内、先于 grads−);
  4. t2 = t1 − diff_offset · inv_num(ChainB:Muls + Sub);
  5. y = (t2 · scale) / s(ChainC:Mul + Div),结果回落到 grads 的 dtype。
  • 输入 dtype:grads/x/y ∈ {FLOAT16, FLOAT32, BFLOAT16} 三者一致;5 个参数张量恒 FLOAT32;
  • 布局:grads/x/y 支持 NCDHW(通道轴 dim1)与 NDHWC(通道轴 dim4)双布局;1D 参数沿通道轴广播到 5D,其余轴为广播轴;
  • 中间精度:FLOAT16 / BFLOAT16 输入在 kernel 内先逐元素提升 FLOAT32 参与全部中间运算,结果回落原 dtype(CAST_RINT,round-half-even);
  • 特殊值:全输入幅值 ≥ 3e38 的极端用例下,f32 中间量溢出 ±Inf 与 f64 golden 分类不一致,按 IEEE 754 传播契约做元素级分类修复(FixVF 的 condA/B/C,逐 lane 向量选择、非控制流分支,对常规数据幂等);
  • tiling:PadAndSqueeze 后有效 rank ≤ 4 → RANK_4 / tilingKey=0;= 5 → RANK_8 / tilingKey=1;多核按 tile 均分,无跨核归并、无 SetAtomicAdd(确定性实现);NDHWC 且 C 满足对齐/容量条件时走 Pass1 寄存器参数快速路径(5 参数按 C 长度驻留寄存器跨 w 复用)。
likedislike
SunSun成员
9 天前 将 LuckySun 设为负责人
CANN-robotCANN-robot成员
8 天前 关闭了 issue
CANN-robotCANN-robot成员
8 天前 添加了label:resolved