文件最后提交记录最后更新时间
23 天前
6 个月前
6 个月前
2 个月前
2 个月前
2 个月前
6 个月前
2 个月前
README

BatchNormV3

产品支持情况

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

功能说明

  • 算子功能:对一个批次的数据做正则化处理,正则化之后生成的数据的统计结果为0均值、1标准差。

  • 计算公式:

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

    E(x)表示均值,Var(x)表示方差,均需要在算子内部计算得到;ε表示一个极小的浮点数,防止分母为0的情况。

参数说明

参数名 输入/输出/属性 描述 数据类型 数据格式
x 输入
  • 进行批量归一化的输入张量,对应公式中的`x`。
  • shape支持4D、5D。
FLOAT32、FLOAT16、BFLOAT16 NCHW/NHWC/NCDHW/NDHWC
weight 输入
  • 进行批量归一化的权重,对应公式中的`γ`。
  • 一个1D张量,shape与输入x的维度C相同,数据类型与x的数据类型支持以下组合:[x: float16, weight: float16/float32], [x: bfloat16, weight: bfloat16/float32], [x: float32, weight: float32]。
FLOAT32、FLOAT16、BFLOAT16 ND
bias 输入
  • 进行批量归一化的偏置值,对应公式中的`β`。
  • 一个1D张量,数据类型和shape与入参weight保持一致。
FLOAT32、FLOAT16、BFLOAT16 ND
running_mean 输入
  • 训练场景:训练期间动量更新前的均值;推理场景:推理期间使用的均值,对应公式中的`E(x)`。
  • 一个1D张量,shape与输入x的维度C相同,数据类型与x的数据类型支持以下组合:[x: float16, running_mean: float16/float32], [x: bfloat16, running_mean: bfloat16/float32], [x: float32, running_mean: float32]。
FLOAT32、FLOAT16、BFLOAT16 ND
running_var 输入 训练场景:训练期间动量更新前的方差;推理场景:推理期间使用的方差,对应公式中的`Var(x)`。 FLOAT32、FLOAT16、BFLOAT16 ND
epsilon 可选属性
  • 添加到方差中的小值以避免除以零,对应公式中的`ε`。
  • 默认值为1e-5。
FLOAT32 -
momentum 可选属性
  • 动量参数,用于更新训练期间的均值和方差。
  • 默认值为0.1。
FLOAT32 -
is_training 可选属性
  • 标记是否训练场景,true表示训练场景,false表示推理场景。
  • 默认值为true。
BOOL -
y 输出
  • 表示批量归一化后的输出结果,对应公式中的`y`。
  • 数据类型、数据格式、shape与输入x保持一致。
FLOAT32、FLOAT16、BFLOAT16 NCHW/NHWC/NCDHW/NDHWC
running_mean 输出
  • 只训练场景输出,训练期间动量更新后的平均值。
  • 数据类型、shape与输入running_mean保持一致。
FLOAT32、FLOAT16、BFLOAT16 ND
running_var 输出
  • 只训练场景输出,训练期间动量更新后的方差。
  • 数据类型、shape与输入running_var保持一致。
FLOAT32、FLOAT16、BFLOAT16 ND
save_mean 输出
  • 只训练场景输出,保存的x均值,对应公式中的`E(x)`。
  • 1D张量。
FLOAT32 ND
save_rstd 输出
  • 只训练场景输出,保存的x方差或者x标准差倒数,分别对应公式中的`Var(x)`、(Var(x) + ε)开平方的倒数。
  • 1D张量。
FLOAT32 ND
  • Atlas A2 训练系列产品/Atlas A2 推理系列产品、Atlas A3 训练系列产品/Atlas A3 推理系列产品:

    • 输入x和输出y的数据格式不支持NHWC、NDHWC。
    • 输出参数save_rstd保存的是x方差。
  • Atlas 训练系列产品:

    • 数据类型:所有的输入和输出不支持BFLOAT16。
    • 数据格式:输入x和输出y不支持NHWC、NDHWC。
    • 输出参数save_rstd保存的是x方差。
  • Atlas 推理系列产品:

    • 数据类型:所有的输入和输出不支持BFLOAT16。
    • 数据格式:输入x和输出y不支持NHWC、NDHWC。
    • 输出参数save_rstd保存的是x方差。
  • Ascend 950PR/Ascend 950DT:

    输出参数save_rstd保存的是x标准差的倒数。

约束说明

Atlas A2 训练系列产品/Atlas A2 推理系列产品、Atlas A3 训练系列产品/Atlas A3 推理系列产品:仅支持训练场景。

调用说明

调用方式 样例代码 说明
aclnn接口 test_aclnn_batch_norm_v3 通过aclnnBatchNorm接口方式调用BatchNormV3算子。
图模式 - 通过算子IR构图方式调用BatchNormV3算子。