Thanks for sending an requirement! Please fill in the following template to help quickly solve your problem.
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 归一化训练链路不完整,需新增本算子补齐。
cann/ops-nn 仓库 norm 目录下新增算子需求
补齐 3D BatchNorm 训练更新链路,与前驱 BN3DTrainingReduce 构成完整的 3D 归一化训练算子对,支撑 3D 卷积等网络训练场景;对齐现有 2D BatchNorm 训练 update 的语义与用法,降低上层框架接入成本。
计算公式:记通道数 C = sum.shape[0],reduce 域元素个数 num = x.size / C,逐通道有:
C = sum.shape[0]
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
multiplier = scale · rsqrt(σ²_biased + ε)
addend = offset − multiplier · μ
y = x · multiplier + addend
mean_out = mean·(1−factor) + batch_mean·factor
variance_out = variance·(1−factor) + σ²_unbiased·factor
num == 1
注:x 支持 fp16/fp32/bf16,统计量恒 fp32;输入格式支持 NCHW/NCDHW/NHWC/NDHWC。
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 + offsetbatch_mean = μ、batch_variance = σ²_biased(有偏,供反向复用)mean_out = factor · μ + (1 − factor) · meanvariance_out = factor · σ²_unbiased + (1 − factor) · variance(σ²_unbiased = num/(num−1) · σ²_biased,Bessel 修正)multiplier = scale · rsqrt(σ²_biased + ε),addend = offset − multiplier · μ;y = x · multiplier + addend完成,中间精度 fp32,y 回 cast 到 x 的 dtype;mean_out = mean·(1−factor) + batch_mean·factor,variance_out = variance·(1−factor) + σ²_unbiased·factor,inplace 写回输入 mean/variance;num == 1时 Bessel 修正分母为 0,显式置无偏 batch variance 为 0,此时 running variance 不更新(保留旧值),batch_variance(有偏)正常输出。