已关闭
[Requirement|需求建议]: 下一代支持BN3DTrainingUpdate #5004
Sun创建于  14 天前关闭于  14 天前
Sun
Sun成员
14 天前 创建

Thanks for sending an requirement! Please fill in the following template to help quickly solve your problem.

Backgroud(背景信息)

BN3DTrainingUpdate 用于 3D BatchNorm 训练流程的"更新"环节:算子接收前驱 BN3DTrainingReduce 产出的逐通道统计量(sum、square_sum),结合 scale、offset 对输入 x 做批归一化得到输出 y;同时计算当前 batch 的 save 统计量(batch_mean、batch_variance)供反向传播复用,并以 factor(等价于 PyTorch 的 momentum)为 EMA 权重更新 running mean 与 running variance(mean_out、variance_out,inplace 写回输入 mean、variance)。

当前 ops-nn 的 norm 目录下已有 BNTrainingUpdate 系列(1D/2D 归一化训练更新链路),缺少 3D BatchNorm 对应的 update 算子,导致 3D 归一化训练链路不完整,需新增本算子补齐。

Origin(信息来源)

cann/ops-nn 仓库 norm 目录下新增算子需求

Benefit / Necessity (价值/作用)

补齐 3D BatchNorm 训练更新链路,与前驱 BN3DTrainingReduce 构成完整的 3D 归一化训练算子对,支撑 3D 卷积等网络训练场景;对齐现有 2D BatchNorm 训练 update 的语义与用法,降低上层框架接入成本。

Design(设计方案)

计算公式:记通道数 C = sum.shape[0],reduce 域元素个数 num = x.size / C,逐通道有:

  • μ = sum / num
  • σ²_biased = square_sum / num − μ²
  • y = (x − μ) / √(σ²_biased + ε) · scale + offset
  • batch_mean = μbatch_variance = σ²_biased(有偏,供反向复用)
  • mean_out = factor · μ + (1 − factor) · mean
  • variance_out = factor · σ²_unbiased + (1 − factor) · varianceσ²_unbiased = num/(num−1) · σ²_biased,Bessel 修正)
  1. 由逐通道统计量预算归一化系数:multiplier = scale · rsqrt(σ²_biased + ε)addend = offset − multiplier · μ
  2. 逐元素归一化以 FMA 形式 y = x · multiplier + addend 完成,中间精度 fp32,y 回 cast 到 x 的 dtype;
  3. EMA 更新:mean_out = mean·(1−factor) + batch_mean·factorvariance_out = variance·(1−factor) + σ²_unbiased·factor,inplace 写回输入 mean/variance;
  4. 边界:num == 1 时 Bessel 修正分母为 0,显式置无偏 batch variance 为 0,此时 running variance 不更新(保留旧值),batch_variance(有偏)正常输出。

注:x 支持 fp16/fp32/bf16,统计量恒 fp32;输入格式支持 NCHW/NCDHW/NHWC/NDHWC。

likedislike
SunSun成员
14 天前 添加了label:requirement
SunSun成员
14 天前 将 LuckySun 设为负责人
CANN-robotCANN-robot成员
14 天前 关闭了 issue
CANN-robotCANN-robot成员
14 天前 添加了label:resolved