返回专辑
·Johan·4 分钟阅读

BatchNorm:训练推理模式

model.train() 与 model.eval() 切换 BatchNorm 的统计量来源——忘了 eval 或 batch=1 时 running stats 与当前 batch 混用,推理结果会随机飘。

BatchNorm:训练推理模式

1. 同一输入,两次推理结果不同

部署分类服务,客户报同一图片连续请求结果不一致——trace 发现 inference 脚本从未调用 model.eval(),BatchNorm2d 仍用当前 batch 的 mean/var。batch=1 时统计量噪声极大,输出随机飘。BN 的行为由 train/eval 模式决定,与 Dropout 一样是部署必查项。

2. train 与 eval 各用什么

PyTorch BatchNorm2dmodel.train() 时用当前 batch 的 mean/var 归一化,并 momentum 更新 running_mean、running_var;model.eval() 时用 running_mean/running_var,不再更新,也不依赖当前 batch 统计。

python
model.train()  # BN 用当前 batch mean/var,更新 running stats
model.eval()   # BN 用 running_mean, running_var,不更新

torch.no_grad() 只关梯度,不切换 BN 模式。validation loop 必须 model.eval(),否则 val metric 带 batch 噪声且与 deploy 不一致。

3. batch size 与 running stats

训练 batch 小(如 2–4)时,batch 统计量方差大,running stats 更新不稳定——微调时常 freeze BN 或换 GroupNorm。推理 batch=1 在 eval mode 下没问题(用 running stats);若在 train mode 下 batch=1,等价于用单样本统计,结果不可信。

SyncBatchNorm 多卡聚合 batch 统计,等效大 batch——多卡训练小 per-GPU batch 时有用。export ONNX 前务必 eval,否则 BN 节点行为可能固化错误。

4. 与 Dropout、LayerNorm 的对比

Dropout:train 随机 drop,eval 关闭。BatchNorm:train 用 batch stats + 更新 running,eval 用 running。LayerNorm/GroupNorm:不依赖 batch 统计,train/eval 行为一致——小 batch 微调时常换 GN 避免 BN 问题。

Checkpoint 保存的 running_mean/var 是模型状态一部分;load 后 eval 依赖这些值。从 train checkpoint 直接 export 而不 eval,是常见 bug。

5. 案例:batch=1 在线推理抖动

某边缘设备 batch=1 推理,结果帧间抖动——未 eval 导致每帧用不同单样本统计。改 eval 后输出稳定。batch=1 推理必须在 eval mode,这是硬约束。

6. 验收

  • inference 脚本含 model.eval(),CI 测确定性。
  • val metric 在 eval mode 下计算。
  • 小 batch 训练时评估 BN 是否 freeze 或换 GN。
  • export 前后同一 input 数值一致(允许 fp 误差)。
  • checkpoint load 后 running stats 随 state_dict 恢复。

7. 落地与衔接

Inference 与 export 脚本在 load 权重后首行 model.eval(),health check 对固定 tensor 连跑两次 assert 数值一致。val/test 循环封装成共用函数,内部强制 eval,避免各脚本 copy-paste 漏切换。小 batch 微调文档写明 freeze BN 或换 GN 的策略,并与 checkpoint 里 running_mean/var 来源对齐。ONNX 导出前后用同一 input 对比输出,允许微小 fp 误差但不可系统性漂移。发版清单把批归一化与随机失活并列为必查项,审查时专门看验证循环是否误开训练模式。

8. 案例复盘

某边缘设备批量为一推理帧间抖动,根因是服务未切评估模式,每帧用单样本批次统计。改评估模式后输出稳定,持续集成增加确定性测试。另项目验证比测试高两点,查为验证循环误开训练模式更新批归一化。现在指标与推理共用评估封装,回归不再靠人工记得切换。该故障已纳入上线前冒烟脚本,减少重复踩坑。

9. 部署契约

BatchNorm 的 train/eval 切换是部署契约。忘了 eval,serving 用错误统计量——比权重错更隐蔽,因为「还能用,只是不稳定」。running stats 与权重同发版归档;域重估统计量时单独 bump 版本。健康检查对固定输入连跑两次,输出非确定性则拒绝流量。导出与 serving 共用同一评估封装,避免训练脚本与推理服务各写一套模式切换。

10. 域偏移与 running stats

源域预训练的 running_mean/var 到目标域可能系统性偏。常见做法:目标域上 train() 若干步但冻结权重、只更新 BN 统计(AdaBN 一类);或小 batch 微调时直接 freeze BN、只训后面头。无论哪种,验收要对比「源域 stats」与「目标域重估 stats」在目标切片上的指标差——差大说明 stats 本身就是域适配旋钮,不能 silently 沿用源域。发版说明写清 stats 来源(源域 / 目标域重估 / freeze),避免两套 checkpoint 混装。这是部署契约的最后一环。

← 全部文章

johan's blog