文件最后提交记录最后更新时间
29 天前
26 天前
26 天前
29 天前
26 天前
29 天前
26 天前
README

BatchNormExt2

产品支持情况

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

功能说明

  • 算子功能:对4D输入张量按通道维做批量归一化(Batch Normalization)。

  • 计算公式:

    y=(x−E(x))Var(x)+ε∗γ+βy = \frac{(x - E(x))}{\sqrt{Var(x) + ε}} * γ + β

    其中,训练模式下 E(x)Var(x) 由当前批次在空间维度上统计得到;推理模式下 E(x)Var(x) 取输入 input_meaninput_varianceε 表示一个极小的浮点数,防止分母为0的情况。

参数说明

参数名 输入/输出/属性 描述 数据类型 数据格式
input_x 输入
  • 进行批量归一化的输入张量,对应公式中的`x`。
  • 一个4D张量,shape为(N, C, H, W)或(N, H, W, C);ND 格式时按 data_format 属性解释 C 轴位置(NCHW→dim1、NHWC→dim3)。
FLOAT16、FLOAT32 NCHW/NHWC/ND
input_scale 输入
  • 进行批量归一化的权重,对应公式中的`γ`。
  • 一个1D张量,shape与输入input_x的通道维C相同。
FLOAT32 ND
input_offset 输入
  • 进行批量归一化的偏置值,对应公式中的`β`。
  • 一个1D张量,shape与入参input_scale保持一致。
FLOAT32 ND
input_mean 可选输入
  • 训练场景:须为空;推理场景:推理期间使用的均值,为必选输入,对应公式中的`E(x)`。
  • 一个1D张量,shape与入参input_scale保持一致。
FLOAT32 ND
input_variance 可选输入
  • 训练场景:须为空;推理场景:推理期间使用的方差,为必选输入,对应公式中的`Var(x)`。
  • 一个1D张量,shape与入参input_scale保持一致。
FLOAT32 ND
epsilon 可选属性
  • 添加到方差中的小值以避免除以零,对应公式中的`ε`。
  • 默认值为1e-4f。
FLOAT32 -
data_format 可选属性
  • 指定输入input_x的数据格式,支持"NHWC"、"NCHW"
  • 默认值为"NHWC"。
STRING -
is_training 可选属性
  • 标记是否训练场景,true表示训练场景,false表示推理场景。
  • 默认值为true。
BOOL -
output_y 输出
  • 表示批量归一化后的输出结果,对应公式中的`y`。
  • 数据类型、数据格式、shape与输入input_x保持一致。
FLOAT16、FLOAT32 NCHW/NHWC/ND
output_mean 输出
  • 训练模式:当前批次的均值(有偏),推理模式:等于输入input_mean。
  • 一个1D张量,shape与入参input_scale保持一致。
FLOAT32 ND
output_variance 输出
  • 训练模式:当前批次的方差(无偏,贝塞尔校正),推理模式:等于输入input_variance。
  • 一个1D张量,shape与入参input_scale保持一致。
FLOAT32 ND
output_reserve_space_1 输出
  • 为梯度计算预留。训练模式:保存的均值(等于output_mean),推理模式:等于输入input_mean。
  • 一个1D张量,shape与入参input_scale保持一致。
FLOAT32 ND
output_reserve_space_2 输出
  • 为梯度计算预留。训练模式:保存的inv_var:1/sqrt(epsilon + variance),用于反向梯度计算中重用,推理模式:等于输入input_variance。
  • 一个1D张量,shape与入参input_scale保持一致。
FLOAT32 ND

约束说明

  • 输入input_x仅支持4D张量,数据格式支持NCHW、NHWC 和 ND(ND 输入按 data_format 属性解释 C 轴位置)。
  • 输入input_x为具体格式(NCHW/NHWC)时必须与 data_format 属性一致;ND 输入不受此限制。
  • 训练模式下,输入input_mean、input_variance必须为空;推理模式下,输入input_mean、input_variance必须提供。
  • 不支持空张量。

调用说明

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