已关闭
[Requirement|需求建议]: SyncBatchNormGatherStats 算子AscendC实现 #4254
chenchen创建于  7月23日关闭于  7月27日
chenchen
chenchen
7月23日 创建

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 合法性校验,并校验属性 momentumeps

tiling 策略:

1)分核策略

按通道维 C 分核。先按 UB 可用空间反推单核一次可处理的 C 元素数 cTileNum,再由
ubOuter = ceil(C / cTileNum)blockFormer = ceil(ubOuter / coreNum) 推出实际使用核数,优先用满核;
核间不能均分时余数块落到前若干核上。分核结果还要过两道约束:

  • 每核 C 切片按 GM data cache line(64 字节)对齐。Atlas 推理系列产品走标量 kernel,用 SetValue 写 GM 后
    按整条 cache line 回写,若某个核分到的输出切片不足一条完整 cache line,多核回写会互相覆盖并在输出中留下 0;
    该不变量在 tiling 内显式断言,避免后续改动误破坏。
  • 小 C 场景放宽换并行度。当按 64 字节口径分核得到的核数不足 4 时,放宽到 16 元素/核;其余情况维持 64 字节口径
    (实测强行切到更多核反而更慢,核数并非越多越好)。

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__ == 200arch20/ 标量实现,否则为向量实现),
两套实现共用同一份 GM 句柄封装与 tiling 解析基类。

Atlas A2 训练系列产品(向量实现):Init 与 Process 两阶段,Process 内为 CopyIn → Compute → CopyOut 三段流水。

  1. CopyIn 用 DataCopyPad[N, cTile] 带 stride 搬入 total_sum / total_square_sum,同时搬入 mean / variance
  2. fp16 输入先 Cast 到 fp32,再用 ReduceSum<Pattern::Reduce::RA> 沿 N 方向归约出 Σ totalSumΣ totalSquareSum
    fp32 大 C 场景改用逐行 Add 直接累加(单行或 16 行 micro-group),以保留每核一条长 C block、避免重复矩阵归约;
  3. Compute 全部在 fp32 上完成 batchMean / batchVar / batchInvstd 与 running 统计量更新;
  4. CopyOut 将结果 Cast 回原 dtype 后用 DataCopyPad 写出四路输出(batchMeanbatchInvstd 与原地更新的
    running mean / running var)。C tile 队列开启 double buffer。

Atlas 推理系列产品(标量实现):该架构不支持 DataCopyPad GM→UB,故全程用 GetValue / SetValue
配合 fp32 标量公式逐通道计算,按 C 分核;四路输出写完后按 32 字节步长做 DataCacheCleanAndInvalid
与 host 侧「每核 ≥ 64 字节」的分核下限共同保证多核标量回写不丢数。

数值保护Σ sampleCount == 0 时倒数置 0;Σ sampleCount <= 1 时无偏方差分母为 0、样本方差本身恒为 0,
此时退化为有偏估计,避免 Inf * 0 = NaN 写入 running var 污染后续推理。

likedislike
chenchen
chenchen
7月23日 评论:

算子交付路径:experimental/norm/sync_batch_norm_gather_stats/,配套 aclnn 两段式接口
aclnnSyncBatchNormGatherStatsGetWorkspaceSize / aclnnSyncBatchNormGatherStats

关联 PR:https://gitcode.com/cann/ops-nn/merge_requests/6149

likedislike
Ffulltower成员
7月23日 将 fullt 设为负责人
fulltower成员
7月23日 评论:

我们会安排评审

likedislike
CANN-robotCANN-robot成员
7月27日 关闭了 issue
CANN-robotCANN-robot成员
7月27日 添加了label:resolved