swa是训练后期对多个检查点权重进行简单平均的后处理式模型集成方法,能稳定提升泛化能力,在resnet、vit等结构上验证准确率通常提升0.3%–1.0%,但需模型已收敛且学习率已退火;启用时须协同主优化器,正确初始化averagedmodel、注册swalr、定期调用update_parameters,并在预测前swap参数及重校准bn统计量。

Stochastic Weight Averaging 是什么,它真能提升泛化?
SWA(torch.optim.swa_utils.AveragedModel)不是训练策略,而是一种后处理式模型集成——它在训练后期以固定间隔采样多个检查点,对权重做简单平均。实测中,它常使验证准确率提升 0.3%–1.0%,尤其在 ResNet、ViT 等结构上效果稳定;但前提是模型已收敛且学习率已退火到较低水平,否则平均会破坏优化轨迹。
如何正确启用 SWA:关键三步不能错
SWA 必须与主优化器协同工作,不能单独调用。典型错误是把 SWALR 当成替代学习率调度器,其实它只负责控制 SWA 阶段的学习率衰减,主优化器仍需正常更新参数。
- 训练前半段(如前 70% epoch)照常训练,不启用 SWA
- 到达切换点后,初始化
AveragedModel并注册SWALR调度器,注意swa_freq参数必须是整数(如每 5 个 epoch 更新一次平均权重) - 每次调用
swa_model.update_parameters(model)前,确保model的state_dict()已更新(即optimizer.step()已执行)
SWA 模型预测时的常见报错与修复
直接用 swa_model 进行推理会出错:RuntimeError: Input type (torch.cuda.FloatTensor) and weight type (torch.float64) should be the same——这是因为 AveragedModel 默认使用 torch.float64 存储平均权重,而模型本身是 float32。
Python 3.14.2是Python编程语言在2025年12月5日发布的稳定版本,属于3.14系列的第二个维护更新。该版本包含了18项修复,重点解决了多进程、数据类及正则表达式等模块的回归问题,并修复了CVE-2025-12084等安全漏洞。此版本标志着自由线程模式(移除GIL)正式获得官方支持,是Python发展的重要里程碑。
- 修复方式:初始化时显式指定
device和dtype,例如AveragedModel(model, device=model.device, dtype=torch.float32) - 预测前务必调用
swa_model.swap_parameters_with_model()切换回平均权重,否则仍用原始模型参数 - 验证集评估必须在
swap_parameters_with_model()后进行,且之后要再调用一次该方法恢复原始参数(否则影响后续训练)
SWA + Dropout 或 BatchNorm 会怎样?
SWA 对 Dropout 无影响(推理时自动关闭),但 BatchNorm 统计量必须重校准——因为平均后的权重分布可能偏移。PyTorch 官方建议在 SWA 阶段结束后,用训练数据跑 1–5 个 epoch 的 torch.nn.BatchNorm2d.train() 模式,更新 running_mean / running_var。
- 不要用验证集或测试集校准,会泄露标签信息
- 校准期间禁用梯度(
torch.no_grad()),只做前向传播 - 若模型含
InstanceNorm或GroupNorm,则无需校准,它们不依赖全局统计量
SWA 的收益很实在,但容易卡在权重类型不匹配或 BN 统计量未更新这两个点上;一旦绕过,就是开箱即用的泛化增强手段。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!










