已关闭
[Requirement|需求建议]: SyncBatchNormGatherStats 算子AscendC实现 #4254
chenchen创建于 7月23日关闭于 7月27日
chenchen
7月23日 评论:
7月23日 评论:
算子交付路径:experimental/norm/sync_batch_norm_gather_stats/,配套 aclnn 两段式接口
aclnnSyncBatchNormGatherStatsGetWorkspaceSize / aclnnSyncBatchNormGatherStats。


7月27日 关闭了 issue
7月27日 添加了label:resolved
Thanks for sending an requirement! Please fill in the following template to help quickly solve your problem.
Backgroud(背景信息)
使用 Ascend C 对 TBE 实现的 SyncBatchNormGatherStats 算子进行重构,实现了 Ascend C 版本 SyncBatchNormGatherStats 算子对 Atlas A2 训练系列产品/Atlas A2 推理系列产品与 Atlas 推理系列产品的适配。
Origin(信息来源)
外部贡献
Benefit / Necessity (价值/作用)
SyncBatchNormGatherStats 算子用于多卡数据并行训练中的同步批归一化:收集所有 device 上的通道特征和、 特征平方和与样本计数,归约出全局批均值与标准差倒数,并按 momentum 更新全局的 running mean / running var。 它是 SyncBatchNorm 在多卡训练下保证与单卡等价统计量的关键环节,广泛用于目标检测、图像分割等 per-card batch 较小、必须跨卡同步 BN 统计量的视觉训练场景。
Design(设计方案)
host 侧设计
shape / dtype 校验:对
total_sum/total_square_sum(2 维[N, C])、sample_count/mean/variance(1 维)逐项做维度、shape 一致性与 dtype 合法性校验,并校验属性
momentum、eps。tiling 策略:
1)分核策略
按通道维 C 分核。先按 UB 可用空间反推单核一次可处理的 C 元素数
cTileNum,再由ubOuter = ceil(C / cTileNum)、blockFormer = ceil(ubOuter / coreNum)推出实际使用核数,优先用满核;核间不能均分时余数块落到前若干核上。分核结果还要过两道约束:
SetValue写 GM 后按整条 cache line 回写,若某个核分到的输出切片不足一条完整 cache line,多核回写会互相覆盖并在输出中留下 0;
该不变量在 tiling 内显式断言,避免后续改动误破坏。
(实测强行切到更多核反而更慢,核数并非越多越好)。
2)单核内切分策略
充分利用 UB 的原则。按 dtype(fp16 需要额外的 fp32 cast 缓冲)、double buffer、输出队列深度与中间量暂存逐项
折算「每个 C 元素占用的 UB 字节数」,据此求解 C 方向的 tile 大小;当 C 很大、单核需要多轮 UB 循环时,改为在 N
方向再切一层(
nTileNum由 UB 预算搜索得到),保证每核仍保持一条尽量长的连续 C block。3)tilingkey 规划策略
只区分 N 维是否全载 UB 两种执行形态:
10001(N full load):N 行统计量整体驻留 UB,kernel 侧用一次矩阵归约完成沿 N 的求和;20001(N not full load):C 较大或 fp32 流式路径,kernel 侧按 N 分块循环累加。dtype 差异由模板参数在 kernel 内静态分支处理,不额外占用 tilingkey。
kernel 侧设计
kernel 入口按架构选择实现头(
__CCE_AICORE__ == 200→arch20/标量实现,否则为向量实现),两套实现共用同一份 GM 句柄封装与 tiling 解析基类。
Atlas A2 训练系列产品(向量实现):Init 与 Process 两阶段,Process 内为 CopyIn → Compute → CopyOut 三段流水。
DataCopyPad按[N, cTile]带 stride 搬入total_sum/total_square_sum,同时搬入mean/variance;Cast到 fp32,再用ReduceSum<Pattern::Reduce::RA>沿 N 方向归约出Σ totalSum、Σ totalSquareSum;fp32 大 C 场景改用逐行
Add直接累加(单行或 16 行 micro-group),以保留每核一条长 C block、避免重复矩阵归约;batchMean/batchVar/batchInvstd与 running 统计量更新;Cast回原 dtype 后用DataCopyPad写出四路输出(batchMean、batchInvstd与原地更新的running mean / running var)。C tile 队列开启 double buffer。
Atlas 推理系列产品(标量实现):该架构不支持
DataCopyPadGM→UB,故全程用GetValue/SetValue配合 fp32 标量公式逐通道计算,按 C 分核;四路输出写完后按 32 字节步长做
DataCacheCleanAndInvalid,与 host 侧「每核 ≥ 64 字节」的分核下限共同保证多核标量回写不丢数。
数值保护:
Σ sampleCount == 0时倒数置 0;Σ sampleCount <= 1时无偏方差分母为 0、样本方差本身恒为 0,此时退化为有偏估计,避免
Inf * 0 = NaN写入 running var 污染后续推理。