已关闭
[Requirement|需求建议]: BatchNormalizationGrad反向传播算子AscendC实现贡献 #2514
peihaobo创建于  5月6日关闭于  7月8日
peihaobo
5月6日 创建

Thanks for sending an requirement! Please fill in the following template to help quickly solve your problem.

Backgroud(背景信息)

使用 AscendC 对 BatchNormalizationGrad 算子进行实现,该算子是 Batch Normalization的反向传播算子。Batch Normalization 是深度学习中最关键的归一化技术之一,前向传播时对每个 channel 内的激活值做减均值除标准差的标准化,再通过可学习的 scale(γ) 和 shift(β) 做线性变换。反向传播时需计算三个梯度:对输入 x 的梯度 ∂L/∂x、对 scale 的梯度 ∂L/∂γ、对 shift 的梯度 ∂L/∂β。数学公式为(设 m=N×H×W 为单 channel 元素数):x̂ = (x - μ)/σ(即 (input - save_mean) × save_invstd);g_bar = Σ(grad_output) / m;g_x̂_bar = Σ(grad_output × x̂) / m;grad_input = γ × σ⁻¹ × (grad_output - g_bar - x̂ × g_x̂_bar);grad_weight = Σ(grad_output × x̂);grad_bias = Σ(grad_output)。该算子的核心特点是 channel 内的跨元素依赖——每个 channel 的所有元素(跨 batch 和所有空间位置)共同贡献 g_bar 和 g_x̂_bar 两个标量统计量,单个元素的梯度依赖全 channel 的全局信息。实现适配了 Atlas A2 训练系列产品/Atlas A2 硬件(Ascend910B4),满足在 Ascend 平台上深度学习模型训练中 BN 反向传播的底层算子需求。

Origin(信息来源)

哈工大算子团队

Benefit / Necessity (价值/作用)

BatchNormalizationGrad 算子采用 channel 级多核并行策略:各核分配互不重叠的 channel 子集,每个 channel 内先做累加扫描得到 g_bar 和 g_x̂_bar 两个标量统计量,再用这两个统计量反算出梯度写回 grad_input。多核模式下各核直写自己负责的 grad_weight/grad_bias(通道天然不重叠,无需核间归约 workspace),整体 workspace 开销为零。算子支持两种空间计算模式:当 numSpatial 较小时采用整空间模式——将单个 batch 所有空间元素一次装入 UB,最大化双缓冲流水线效率;当 numSpatial 超过 UB 容量时自动降级为空间分块模式——按 spatialTileSize 分多遍扫描,保证超大空间场景(如大分辨率特征图)下的正确运行。累加阶段采用 Kahan 补偿求和减少浮点累积误差,ReduceSum 后通过 V_S 硬件同步事件确保流水线数据就绪。Batch Normalization 在深度学习中的地位举足轻重:它通过减少内部协变量偏移(Internal Covariate Shift)使训练更深网络成为可能、允许更大的学习率加速收敛、降低权重初始化敏感性、引入 mini-batch 噪声起到正则化效果(常可替代 Dropout)、平滑损失曲面改善优化器收敛行为。BN 反向传播的正确实现直接关系到整个网络的梯度质量。该算子使用 6 个 GB→UB 输入(grad_output/input/weight/bias/save_mean/save_invstd)、3 个输出(grad_input/grad_weight/grad_bias)、支持 float32、float16 和 bfloat16 三种数据类型,满足实际大模型训练中 Batch Normalization 反向传播的严苛正确性和性能需求。

Design(设计方案)

host 侧设计:

1)分核策略

采用基于 channel 的多核分配策略(各 channel 相互独立,但 channel 内部存在跨 batch/spatial 的归约统计量):

  • usedCoreNum = min(hardwareCoreNum, numChannels) — 仅在 numChannels > 1 时启用多核(单 channel 无法拆分);
  • channelsPerCore = numChannels / usedCoreNum — 基准 channel 数;
  • tailChannels = numChannels % usedCoreNum — 前 tailChannels 个核各多分配一个 channel;
  • 特殊情况(numChannels == 1):单核模式,无需 SyncAll。
2)单 core 内切分策略

使用双缓冲 TQue(BUFFER_NUM=2),以 batch/spatial 为数据单位。分两步决策:

Step 1:整空间模式可行性检查

  • batchesPerTile:由 UB 容量约束确定,UB 固定开销(6 个 float TBuf × paddedSpatial + 3 TQue × 2 × paddedSpatial × typeLength + scalarBuf + accumUb + 安全余量 2048B)后剩余空间按 per-batch UB 增量分配;
  • 对齐条件:batchStrideAligned 时允许多 batch 合并,否则小 S(≤64)允许多 batch、大 S 退回单 batch;
  • 整空间 UB:CalcWholeSpatialUb(paddedSpatial, batchesPerTile, typeLength) = tqueUb(6×tileElems×typeLength) + tbufUb(6×tileElems×4B) + reduceUb(CalcReduceUb) + scalarUb(32B);

Step 2:空间分块判定(整空间模式 UB 超限时触发)

  • spatialSplitMode=1 时 batchesPerTile=1,分 batch 独立处理;
  • tileCols 二分搜索(SearchSpatialTileSize):在 UB 容量(含 batchGBuf/batchGXhatBuf 和安全余量)约束下最大化 spatialTileSize,对齐到 elementsPerBlock;

UB 缓冲区布局

  • 3 TQue × 2 双缓冲:inQueueGradOut、inQueueInput、outQueueGradInput(原生类型 T)
  • 6 TBuf float:gradOutFloatBuf、inputFloatBuf、xhatBuf、tmpBuf、reduceBuf、scalarBuf(8×4B)
  • 空间分块模式加:reduceWsBuf(按 CalcReduceUb 公式)、batchGBuf + batchGXhatBuf(各 batchesPerTile × 4B)
  • 多核模式加:accumBuf(2 × alignedCPC × 4B)
3)tilingkey 规划策略

采用 ASCENDC_TPL 模板参数机制(schMode 位宽=2):

  • tilingKey=0:float32(DT_FLOAT),实例化 KernelBatchNormalizationGrad
  • tilingKey=1:float16(DT_FLOAT16),实例化 KernelBatchNormalizationGrad
  • tilingKey=2:bfloat16(DT_BF16),实例化 KernelBatchNormalizationGrad<bfloat16_t>

Host 侧根据 grad_output/input 的数据类型设置 tilingKey。epsilon 作为 attr 属性(默认 1e-5)通过 tiling 数据传递到 kernel 侧。Workspace 为零(多核各核直写互不重叠的 channel,无需核间归约)。

kernel 侧设计:

分为 Init 和 Process 两个阶段。float32/float16 共用通用模板类,bfloat16 提供独立的全特化实现。

1)通用模板类 KernelBatchNormalizationGrad(float32 / float16)

Init 阶段 — 完成资源初始化:

  • 读取 Tiling 参数:numChannels/numBatches/numSpatial/mFloat/coreNum/spatialSplitMode/spatialTileSize/spatialTileNum/batchesPerTile/batchGroups/paddedSpatial 等;
  • 计算本核 channel 范围,GM 地址绑定 8 个全局 Tensor;
  • Pipe/TQue/TBuf 初始化:整空间模式按 paddedSpatial 分配 Queue + bufBS 分配 TBuf,空间分块模式按 tileAligned 分配。

Process 阶段 — 根据 spatialSplitMode 和 coreNum 走四种路径之一:

模式 coreNum 函数
整空间 单核 1 ProcessSingleCore → AccumulateChannel + WriteGradInput
整空间 多核 >1 ProcessMultiCore → AccumulateChannel(各核) + WriteGradOutputs(直写) + SyncAll + WriteGradInput
分块 单核 1 ProcessSingleCoreTiled → AccumulateChannelTiled + WriteGradInputTiled
分块 多核 >1 ProcessMultiCoreTiled → AccumulateChannelTiled(各核) + WriteGradOutputs + SyncAll + WriteGradInputTiled

整空间模式 AccumulateChannel 内部

步骤 操作 API 说明
Load DataCopyPad GM→UB DataCopyPad 批量搬运 grad_output + input,双缓冲流水
Cast T→float Cast/Adds float16→float32,float32 直接用 Adds(x,0)
x̂ = (input-mean)×invstd Adds + Muls 标量广播,逐元素计算归一化值
t1 t1 = grad_output × x̂ Mul 逐元素乘法
sumG sumG += ReduceSum(grad_output) ReduceSum Kahan 补偿累加(cG 补偿项),WaitVectorScalarSync 确保标量就绪
sumGXhat sumGXhat += ReduceSum(t1) ReduceSum Kahan 补偿累加(cX 补偿项),WaitVectorScalarSync 确保标量就绪

整空间模式 WriteGradInput 内部

步骤 操作 说明
Load 再次搬运 grad_output + input 与 AccumulateChannel 同
Compute γ×σ⁻¹×(grad_output - gBar - x̂×gXhatBar) Muls(x̂,gXhatBar)→Adds(+gBar)→Sub(grad-t1)→Muls(×weightInvstd)
CastBack float→T Cast/Adds,写回 grad_input

空间分块模式:整空间搬一个完整 spatial 片,空间分块每 tile 只搬 tileAligned 个 spatial 元素(通过 LoadSpatialTile)。AccumulateChannelTiled 和 WriteGradInputTiled 额外循环 spatialTileNum 次。

多核 WriteGradOutputs:各核将 accumBuf 中的 gXhatBar/gBar 通过 ScalarCast(float→T)后 DataCopyPad 直写到 gradWeightGm/gradBiasGm 对应的 channel 偏移(通道不重叠,无需核间同步)。

2)bfloat16 全特化 KernelBatchNormalizationGrad<bfloat16_t>

bfloat16 特化与通用模板差异:

差异点 通用模板 bf16 特化
ScalarCast 入 static_cast<float>(t.GetValue()) ToFloat(t.GetValue())
ScalarCast 出 static_cast<T>(floatVal) ToBfloat16(floatVal)
Cast 入 Cast(dst,src,CAST_NONE) / Adds(×0) Cast(dst,src,CAST_NONE)
Cast 出 Cast(dst,src,CAST_NONE) / Adds(×0) Cast(dst,src,CAST_RINT)
Queue 类型 T(2B/4B) bfloat16_t(2B)
同步函数 WaitVectorScalarSync Bf16WaitVectorScalarSync(独立实例化)
文件位置 batch_normalization_grad.h batch_normalization_grad_bf16.h(#include 引入)

GM 地址计算:GetGmIdx(b, ch, offset) = b×numChannels×numSpatial + ch×numSpatial + offset,按 ND 格式自然排列。

尾块处理:

  • 整空间末 batch 组:最后一个 batch 组的实际 batch 数为 tailBatches(= numBatches % batchesPerTile);
  • 空间分块末 tile:最后一个 spatial tile 的实际元素数为 spatialTailSize,DataCopyPad 的 padElems 据此精确调节;
  • channel 尾块:前 tailChannels 个核各多处理一个 channel,大小核 channel 边界由 channelStart/channelEnd 界定。

数据类型支持:float32(DT_FLOAT)、float16(DT_FLOAT16)、bfloat16(DT_BF16)

likedislike
Ppeihaobo
5月6日 修改了issue 的描述
peihaobo
5月6日 评论:

/assign

likedislike
CANN-robotCANN-robot成员
5月6日 将 xiaopei-1 设为负责人
Ppeihaobo
5月18日 关联了pull request:feat:Add BatchNormalizationGrad operator implementation
Ppeihaobo
5月25日 修改了issue 的描述
Ppeihaobo
7月2日 修改了issue 的描述
CANN-robotCANN-robot成员
7月8日 关闭了 issue
CANN-robotCANN-robot成员
7月8日 添加了label:resolved