文件最后提交记录最后更新时间
14 天前
15 天前
15 天前
15 天前
15 天前
15 天前
15 天前
3 个月前
14 天前
README

SyncBatchNormBackwardElemt

产品支持情况

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

功能说明

  • 算子功能:SyncBatchNormBackwardElemt算子用于计算输入张量的元素级梯度,以便在反向传播过程中更新模型参数。

  • 计算公式:

    gradInput=((gradOut−meanDy)−(input−mean)∗(invstd2∗meanDyXmu))∗invstd∗weightgradInput = ((gradOut - meanDy) - (input - mean) * (invstd^2 * meanDyXmu)) * invstd * weight

参数说明

参数名 输入/输出 描述 数据类型 数据格式
grad_output 输入 表示正向输出的微分,对应公式中的`gradOut`。 FLOAT32、FLOAT16、BFLOAT16 ND
save_input 输入 表示进行BatchNorm计算的输入,对应公式中的`input`。 FLOAT32、FLOAT16、BFLOAT16 ND
mean 输入 表示输入数据均值,对应公式中的`mean`。 FLOAT32、FLOAT16、BFLOAT16 ND
invstd 输入 表示输入数据标准差倒数,对应公式中的`invstd`。 FLOAT32、FLOAT16、BFLOAT16 ND
weight 输入 表示权重Tensor,对应公式中的`weight`。 FLOAT32、FLOAT16、BFLOAT16 ND
mean_dy 输入 表示输出梯度的样本均值和的平均值,对应公式中的`meanDy`。 FLOAT32、FLOAT16、BFLOAT16 ND
mean_dy_xmu 输入 表示样本均值和与输入梯度乘积的平均值,对应公式中的`meanDyXmu`。 FLOAT32、FLOAT16、BFLOAT16 ND
grad_input 输出 表示输入Tensor的梯度,对应公式中的`gradInput`。 FLOAT32、FLOAT16、BFLOAT16 ND

约束说明

参数grad_output、save_input、mean、invstd、weight、mean_dy、mean_dy_xmu、grad_input支持的组合如下所示:

  • Ascend 950PR/Ascend 950DT、Atlas A3 训练系列产品/Atlas A3 推理系列产品、Atlas A2 训练系列产品/Atlas A2 推理系列产品:

    grad_output save_input mean invstd weight mean_dy mean_dy_xmu grad_input
    FLOAT16 FLOAT16 FLOAT16 FLOAT16 FLOAT16 FLOAT16 FLOAT16 FLOAT16
    BFLOAT16 BFLOAT16 BFLOAT16 BFLOAT16 BFLOAT16 BFLOAT16 BFLOAT16 BFLOAT16
    FLOAT32 FLOAT32 FLOAT32 FLOAT32 FLOAT32 FLOAT32 FLOAT32 FLOAT32
    FLOAT16 FLOAT16 FLOAT32 FLOAT32 FLOAT32 FLOAT32 FLOAT32 FLOAT16
  • Atlas 训练系列产品:

    grad_output save_input mean invstd weight mean_dy mean_dy_xmu grad_input
    FLOAT16 FLOAT16 FLOAT16 FLOAT16 FLOAT16 FLOAT16 FLOAT16 FLOAT16
    FLOAT32 FLOAT32 FLOAT32 FLOAT32 FLOAT32 FLOAT32 FLOAT32 FLOAT32

调用说明

调用方式 样例代码 说明
aclnn接口 test_aclnn_batch_norm_elemt_backward 通过aclnnBatchNormElemtBackward接口方式调用SyncBatchNormBackwardElemt算子。
图模式 test_geir_sync_batch_norm_backward_elemt 通过算子IR构图方式调用SyncBatchNormBackwardElemt算子。