文件最后提交记录最后更新时间
5 个月前
5 个月前
5 个月前
4 个月前
3 个月前
3 个月前
3 个月前
5 个月前
5 个月前
README

SyncBatchNormBackwardReduce

产品支持情况

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

功能说明

  • 算子功能:SyncBatchNormBackwardReduce用于反向传播过程中计算BatchNorm操作的所需的权重梯度gradWeight和中间量sumDyXmu。

  • 计算公式:

sumDyXmu=sumDyDxPad−sumDy∗meansumDyXmu = {sumDyDxPad} - {sumDy} * {mean}

gradWeight=(sumDyDxPad−sumDy∗mean)∗invertStdgradWeight = ({sumDyDxPad} - {sumDy} * {mean}) * invertStd

参数说明

参数名 输入/输出 描述 数据类型 数据格式
sum_dy 输入 表示正向输出梯度的累加和,对应公式中的`sumDy`。 FLOAT32、FLOAT16、BFLOAT16 ND
sum_dy_dx_pad 输入 对应公式中的`sumDyDxPad`。 FLOAT32、FLOAT16、BFLOAT16 ND
mean 输入 表示输入数据均值,对应公式中的`mean`。 FLOAT32、FLOAT16、BFLOAT16 ND
invert_std 输入 表示输入数据标准差倒数,对应公式中的`invertStd`。 FLOAT32、FLOAT16、BFLOAT16 ND
sum_dy_xmu 输出 表示正向输出梯度与输入中心化后数据乘积之和,对应公式中的`sumDyXmu`。 FLOAT32、FLOAT16、BFLOAT16 ND
y 输出 表示缩放参数的梯度,对应公式中的`gradWeight`。 FLOAT32、FLOAT16、BFLOAT16 ND

约束说明

调用说明

调用方式 样例代码 说明
aclnn接口 test_aclnn_batch_norm_reduce_backward 通过aclnnBatchNormReduceBackward接口方式调用SyncBatchNormBackwardReduce算子。
图模式 test_geir_sync_batch_norm_backward_reduce 通过算子IR构图方式调用SyncBatchNormBackwardReduce算子。