☰
Composer 中的 EMA 指数移动平均算法:原理、超参数与实战接入指南
2026/10/12 3:03:30 网站建设 项目流程
  • 深度学习
  • 分布式训练
  • 模型优化

【免费下载链接】composer

Supercharge Your Model Training

项目地址:https://gitcode.com/gh_mirrors/com/composer
点击查看免费下载

导读

本文以 MosaicML Composer 开源仓库中 composer/algorithms/ema/README.md 为骨架,系统讲解Exponential Moving Average(EMA)模型参数指数移动平均算法。EMA 在训练过程中维护一份对模型参数做指数加权平均的副本,并用这份平均参数进行模型评估,通常能带来更平滑的验证指标与更优的泛化性能。读完本文,你将掌握 EMA 的数学原理与平滑系数换算、Composer 功能接口与 Trainer 两种接入方式、四个关键超参数的语义与推荐取值,以及其内存、计算、评估与检查点保存的注意事项,并能直接从仓库源码层面理解其事件驱动的实现机制。

EMA 是什么:为什么训练中要给权重做平均

训练深度模型时,最后若干次迭代的权重往往落在损失曲面的一个波动区域,单次权重快照得到的验证指标噪声较大,泛化能力也不稳定。EMA 的思路是:在训练过程中持续维护一组指数加权移动平均权重,让早期与近期的参数信息按照指数衰减的方式融合,从而逼近一个更"居中"、更平滑的解。

从 composer/algorithms/ema/README.md 的定义看,Composer 中 EMA 的核心特征有三点:

  • 维护一份平均权重副本:每次迭代用最新训练权重更新平均权重;
  • 用平均权重做评估:训练权重只用于前向/反向与优化器更新,评估(以及默认的检查点保存)使用平均权重;
  • 带来更平滑的验证曲线:由于平均权重的变化是渐进连续的,训练过程中的验证指标通常更平滑、噪声更小,有时还能提升最终泛化能力。

该方法在 Composer 中被归类为cv(计算机视觉)与nlp(自然语言处理)两个领域通用的训练优化算法(见 metadata.json),并不限定于某个特定模型架构。

数学原理:平滑系数、半衰期与权重更新公式

权重更新公式

在 ema.py 中,compute_ema函数按如下公式就地更新平均权重:

W_ema^(t+1) = smoothing × W_ema^(t) + (1 - smoothing) × W_model^(t)

其中W_model是当前训练权重,W_ema是维护的平均权重。smoothing越大,历史信息保留得越多、更新越"缓慢";smoothing越小,越倾向于快速跟随最新权重。

实现时,该函数会遍历模型的named_parameters()与named_buffers(),用copy_在torch.no_grad()下就地更新ema_model中同名的参数与缓冲区(Buffer)。也就是说平均的不仅是可训练参数,也包含 BatchNorm 等模块的统计缓冲区。

半衰期与平滑系数的换算

half_life(半衰期)指平均中一项旧信息"衰减一半"所需的时间步数,它与smoothing之间的换算关系在源码中给出:

t_1/2 = -log(2) / log(smoothing) smoothing = exp[-log(2) / t_1/2]

Composer 的EMA算法类在初始化时会依据传入的half_life与update_interval自动计算平滑系数(见 ema.py):

self.smoothing = 2**(-(update_interval.value / half_life.value))

这个公式可以看作对exp[-log(2) × (update_interval / half_life)]的等价写法——因为每次更新的间隔为update_interval,所以实际衰减量与"间隔占半衰期的比例"直接相关。这一点在 tests/algorithms/test_ema.py 中通过np.exp(-np.log(2) * (update_interval.value / half_life.value))与algorithm.smoothing的比对进行了验证。

例如half_life='1000ba'、update_interval='1ba'时,每次更新对应的平滑系数约为2^(-1/1000) ≈ 0.9993,意味着单次迭代后旧信息保留约 99.93%。

两种接入方式:Functional API 与 Composer Trainer

方式一:Functional 接口(compute_ema)

不依赖 Trainer、希望在自建训练循环中手动接入 EMA 时,使用composer.functional下的compute_ema。README 给出了完整示例骨架:

import copy import composer.functional as cf def training_loop(model, train_loader): opt = torch.optim.Adam(model.parameters()) loss_fn = F.cross_entropy ema_model = copy.deepcopy(model) # 深拷贝一份作为平均权重的载体 model.train() for epoch in range(num_epochs): for X, y in train_loader: y_hat = model(X) loss = loss_fn(y_hat, y) loss.backward() opt.step() opt.zero_grad() cf.compute_ema(model, ema_model, smoothing=0.99) # 每步更新平均权重

要点:

  • ema_model必须预先用copy.deepcopy(model)初始化,且与model结构一致;
  • smoothing必须落在开区间(0, 1)内,默认0.99;
  • 每步迭代在opt.step()之后调用一次compute_ema(model, ema_model, smoothing=0.99),ema_model被就地更新;
  • 评估阶段务必使用ema_model(而非训练模型)进行推理,才能获得平均权重的泛化收益。

除torch.nn.Module外,compute_ema也接受EMAParameters对象(内部存放参数/缓冲区字典的容器),传入其他类型时会抛出ValueError('ema_model must be a torch.nn.Module or EMAParameters')。

方式二:Composer Trainer 算法(EMA)

在 Trainer 模式下,只需实例化EMA算法并放进algorithms列表,Trainer 会在训练循环的适当时机自动完成初始化、更新、评估切换与检查点保存,无需手写任何更新逻辑:

from composer.algorithms import EMA from composer.trainer import Trainer ema = EMA(half_life='50ba') trainer = Trainer(model=model, train_dataloader=train_dataloader, max_duration='1ep', algorithms=[ema]) trainer.fit() model = ema.ema_model

这里的ema.ema_model是EMAParameters实例,用于在训练结束后访问/导出平均权重(也可通过下文介绍的get_ema_model将其写回任意模型)。

两种方式如何选择

  • 需要完全控制训练循环(如自研框架、研究性代码)时选 Functional 接口;
  • 使用 Composer Trainer 时推荐算法类方式,超参数校验、事件调度、评估与检查点切换均由框架自动处理,代码量最少且不易出错。

超参数详解与推荐取值

Trainer 实现中的EMA构造函数签名与默认值如下(见 ema.py):

EMA(half_life='1000ba', smoothing=None, ema_start='0.0dur', update_interval=None)
参数含义默认值说明
half_life平均中各项的半衰期,越长旧信息保留越久,越短旧信息越快被丢弃'1000ba'时间字符串,整数取值,单位仅支持'ba'(batch)与'ep'(epoch);0表示不平均,无穷大表示不更新
update_interval两次更新平均权重之间的间隔,越长更新越稀疏None未指定时:使用half_life则默认1个half_life单位;使用smoothing则默认'1ba'。单位必须与half_life一致
ema_startEMA 开始生效前已完成训练量'0.0dur'支持'dur'(训练总时长比例)、'ba'、'ep'三种单位;'0.0dur'表示从训练一开始就启用
smoothing旧观察的保留系数,须在(0, 1)内None与half_life二选一,不能同时指定;指定后不再随update_interval自动调整

关于时间字符串:Composer 的Time.from_timestring(见 core/time.py)支持"数字 + 单位缩写"的写法,如'5ep'、'1000ba'、'0.5dur';除dur外的单位要求整数取值。

推荐的起始配置

README 给出的典型实践是:

  • half_life='1000ba'(1000 个 batch 的半衰期)作为起始值;
  • update_interval可留空(自动取'1ba'),或设置为更大的值如'10ba'以降低每次更新的开销、提升训练速度;
  • 更短的更新间隔通常带来更好的泛化性能,但会增加少量运行时间。实践中只要half_life远大于update_interval,拉大update_interval对泛化性能的影响很小。

直接使用smoothing(兼容其他实现)

为了与 PyTorch、TensorFlow 等其他生态中的 EMA 实现对齐,Composer 也允许直接指定smoothing:

ema = EMA(half_life=None, smoothing=0.99, update_interval='1ba')

此时half_life必须显式传None,smoothing直接作为更新系数。需要注意:使用smoothing时该值不会随update_interval改变而重新换算,因此修改update_interval会改变平均的时间尺度语义(相当于改变了实际半衰期)。

参数校验规则(来自源码)

EMA.__init__对参数做了严格校验(见 ema.py):

  • half_life与smoothing都未指定时抛出ValueError(二选一必须满足一个);
  • 两者同时指定时抛出ValueError;
  • half_life与update_interval的单位不一致时抛出ValueError;
  • update_interval只允许BATCH或EPOCH单位。

源码级原理:EMA 是如何挂在训练循环上的

事件驱动与状态机

EMA继承自composer.core.Algorithm,通过match(event, state)与apply(event, state, logger)接入 Composer 的事件系统(Event)。从 ema.py 可以看到其事件绑定关系:

  • 初始化与参数搬移:FIT_START、PREDICT_START、EVAL_START时调用move_params_to_device,确保从检查点恢复或设备变化后平均参数落在正确设备上(例如多卡/FSDP 场景);
  • 权重交换时机:BATCH_START、EVAL_START、EVAL_END时在训练权重与平均权重之间swap_params切换;
  • 更新时机:update_interval单位为'ba'时在BATCH_END更新,单位为'ep'时在EPOCH_END更新;更新满足"当前时间步是update_interval的整数倍"这一条件(见 ema.py);
  • 检查点时机:BATCH_CHECKPOINT、EPOCH_CHECKPOINT时若存在CheckpointSaver且到达保存间隔,则触发权重切换,保证保存的是平均权重。

三个内部标志

EMA维护三个序列化状态标志(ema_model、ema_weights_active、ema_started,见 ema.py):

  • ema_started:EMA 是否已启动(由ema_start阈值触发),启动时通过EMAParameters(state.model)从训练模型克隆出平均参数与缓冲区的初始副本;
  • ema_weights_active:当前state.model中装载的到底是平均权重还是训练权重,评估与检查点保存前切换为平均权重,训练(BATCH_START)时切回训练权重。

FSDP 兼容性

对于使用 Fully Sharded Data Parallel(FSDP)的模型,平均参数是分片的,直接操作param.data并不可行。源码为此提供了get_model_context_manager(见 ema.py):当检测到模型是 FSDP 模型时,会在model.module.summon_full_params(...)上下文内执行参数拷贝/交换,EMAParameters.swap_params与transfer_ema_params也统一使用copy_而非裸数据访问(源码注释明确指出"raw data access (eg .data) doesn't work with FSDP")。

状态保存与版本兼容

EMA.state_dict将平均权重以named_parameters_dict与named_buffers_dict字典形式序列化;ensure_compatible_state_dict兼容 Composer 0.13.0 之前同时保存training_model与ema_model两份权重的旧格式检查点,自动将其重写为新格式(见 ema.py)。这意味着老版本训练产出的 EMA 检查点可以直接被新版加载。

评估与检查点:该用哪一份权重

这是实战中最容易踩坑的一点,README 专门用警示块强调了三条规则:

  1. 评估必须用平均权重:Functional 实现中应使用ema_model做推理;Trainer 实现中训练结束后通过model = ema.get_ema_model(model)把平均权重写回传入的模型(若 Composer 模型当前已装载平均权重则无需再写)。
  2. 训练权重可随时取回:通过model = ema.get_training_model(model)恢复未应用 EMA 的训练权重,用于继续训练或对照实验。
  3. 默认检查点保存平均权重:通过CheckpointSaver回调或 Trainer 参数保存检查点时,默认保存的是 EMA 模型权重;唯一例外是显式调用trainer.save_checkpoint(),此时保存的是训练权重并记为state.model。

对应的两个方法定义在 ema.py:get_ema_model在ema_weights_active == True(平均权重已在模型中)时抛错,get_training_model在ema_weights_active == False时抛错,避免重复覆盖造成权重错乱。

成本与注意事项

内存开销

EMA 需要额外保存一份与模型可训练参数 + 缓冲区等大小的权重副本,因此会增大设备内存占用。但注意,这份额外内存只相当于"一份模型参数",激活值(activations)与优化器状态不会被复制,所以相对训练整体的内存占用而言,额外开销通常很小(README 原话:"the extra memory used is small relative to the total amount of memory used")。

计算开销

每次更新都要做一次smoothing × 旧值 + (1 - smoothing) × 新值的逐参数融合,带来少量额外计算与轻微减速。降低该开销的方法是拉大update_interval(如从'1ba'改为'10ba'),让平均计算更稀疏。如前所述,只要half_life远大于update_interval,此举对最终泛化性能影响很小。

与其他平均方法的组合

模型平均类方法(model-averaging methods)一般不推荐叠加使用。README 明确建议:在EMA 与 SWA(Stochastic Weight Averaging)中二选一,不要同时使用。

仓库中的 SWA 实现(见 swa.py)同样维护一份平均权重副本,与 EMA 机制重叠,叠加既增内存又不带来额外收益。

经验结论(来自方法卡片/README)

  • ✅改善质量与训练速度的权衡:实验表明 EMA 能改善训练速度与最终模型质量之间的可达成权衡,官方推荐在卷积网络训练中使用 EMA;
  • ✅验证指标更平滑:只要评估指标在训练过程中周期性计算,EMA 平均权重通常会让这些指标更平滑、噪声更小。

质量验证:测试如何保证 EMA 行为正确

仓库的 tests/algorithms/test_ema.py 对 EMA 的数学行为做了系统验证,可以作为你接入后自测的参考:

  • test_ema:对SimpleConvModel、SimpleTransformerClassifier、Tiny BERT 三类模型,在smoothing ∈ {0, 0.5, 0.99, 1}下调用compute_ema,逐参数校验new = original × smoothing + (1 - smoothing) × param完全成立(参数与缓冲区都验证);
  • test_ema_algorithm:分别覆盖half_life='10ba' + update_interval='1ba'、half_life='1ep' + update_interval='1ep'、smoothing=0.999 + update_interval='1ba'三组配置,验证:自动换算出的smoothing与理论值一致、BATCH_END/EPOCH_END时平均权重更新正确、EVAL_START后state.model被替换为平均权重、EVAL_END后恢复训练权重。

这套测试同时印证了本文前面关于公式换算与"评估/训练权重自动切换"的所有描述。

小结

EMA 是一种低开销、易接入、普适性强的模型平均技术。在 Composer 中,你可以用一行cf.compute_ema(model, ema_model, smoothing=0.99)在自建循环中手动启用,也可以用EMA(half_life='1000ba')交给 Trainer 全自动管理。其核心控制点集中在四个超参数——half_life/smoothing(平均时间尺度)、update_interval(更新频率)、ema_start(启动时机),配合事件系统自动完成权重交换、设备搬移与检查点保存。接入时请牢记三条铁律:评估用平均权重、默认检查点存平均权重、不与 SWA 同用。

相关参考

  • 方法卡片:docs/source/method_cards/ema.md(与 README 内容一致)
  • 核心实现:composer/algorithms/ema/ema.py(EMA类、compute_ema、EMAParameters)
  • 算法元数据:composer/algorithms/ema/metadata.json
  • 测试用例:tests/algorithms/test_ema.py
  • 时间字符串解析:composer/core/time.py
  • SWA(不建议与 EMA 同用):composer/algorithms/swa/swa.py
  • 算法注册导出:composer/algorithms/init.py
  • 深度学习
  • 分布式训练
  • 模型优化

【免费下载链接】composer

Supercharge Your Model Training

项目地址:https://gitcode.com/gh_mirrors/com/composer
点击查看免费下载

相关推荐

上一篇:如何快速搭建个人专属的影视聚合播放站
下一篇:llamafile 项目使用教程

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询