尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

深度学习模型训练中的指数移动平均(EMA)原理与PyTorch实现

深度学习模型训练中的指数移动平均(EMA)原理与PyTorch实现 1. 指数移动平均EMA是什么以及为什么我们需要它如果你在训练深度学习模型尤其是像YOLOv5这类目标检测网络时经常会在代码里看到一个叫ema.py的文件或者在训练日志里看到EMA相关的参数更新。这个EMA全称是指数移动平均乍一听像是金融时间序列分析里的东西怎么就跑进深度学习里了简单来说EMA在深度学习中扮演着一个“影子模型”的角色它不直接参与梯度下降而是默默地跟在你的主模型后面用更平滑、更“保守”的方式更新自己的参数。这个影子模型的性能往往比那个上蹿下跳、正在经历剧烈梯度更新的主模型要更稳定、更鲁棒尤其是在模型训练的中后期和最终评估时。为什么主模型需要这么一个“影子”这得从深度学习的训练动力学说起。我们用的随机梯度下降SGD或其变体如Adam每一步更新都基于一个小批量mini-batch的数据。这个小批量只是全体数据的一个微小采样不可避免地带有噪声。这就导致模型的参数在最优值附近剧烈震荡就像一个人蒙着眼睛在山顶找最低点每一步都跌跌撞撞。虽然从长期平均来看他确实在往山下走但任何单个时间点他所在的位置即模型当前参数都可能不是那个最好的位置。EMA做的就是一件事给这个跌跌撞撞的路径做一个平滑。它通过赋予近期参数更高的权重但又不完全遗忘历史计算出一个移动平均值。这个平均值过滤掉了短期噪声更能反映参数变化的长期趋势因此通常能获得一个更泛化、测试性能更好的模型版本。在YOLOv5、YOLOv8等流行的目标检测框架中EMA是默认开启的。你会发现最终保存的best.pt和last.pt文件旁边往往还有一个best_ema.pt。在验证集上表现最好的经常就是这个EMA模型。所以理解并善用EMA是提升模型实际部署性能的一个简单却有效的技巧。2. EMA的数学原理与核心参数解析2.1 从简单平均到指数移动平均要理解EMA我们先从最简单的概念——算术平均开始。假设我们有一系列模型参数值 ( \theta_1, \theta_2, ..., \theta_t )其算术平均就是 ( \bar{\theta}t \frac{1}{t} \sum{i1}^{t} \theta_i )。这意味着所有历史值权重相同。但在模型训练中我们显然更关心最近的参数因为模型在持续学习进化很久以前的参数可能已经过时了。于是就有了移动平均Moving Average, MA它只考虑最近N个值( MA_t \frac{1}{N} \sum_{it-N1}^{t} \theta_i )。这比算术平均更关注近期但有个问题它给最近N个值赋予了相同的权重而对第N1步之前的数据权重直接降为0这个变化是“硬”的不连续。指数移动平均EMA则提供了一种优雅的平滑过渡。它的公式是 [ \text{EMA}t \beta \cdot \text{EMA}{t-1} (1 - \beta) \cdot \theta_t ] 其中( \text{EMA}_t ) 是当前时刻第t步的EMA值。( \theta_t ) 是当前时刻模型参数的实际值。( \beta ) 是一个介于0和1之间的衰减率decay rate它是EMA的灵魂参数。( \text{EMA}_{0} ) 通常初始化为 ( \theta_0 ) 或 0。这个递归公式的妙处在于它等价于给所有历史参数 ( \theta_i ) 分配了一个随时间指数衰减的权重。把公式展开 [ \text{EMA}t (1-\beta)\theta_t \beta(1-\beta)\theta{t-1} \beta^2(1-\beta)\theta_{t-2} ... \beta^{t-1}(1-\beta)\theta_1 \beta^t \text{EMA}0 ] 可以看到距离现在第 ( k ) 步的参数 ( \theta{t-k} ) 的权重是 ( (1-\beta)\beta^k )。由于 ( \beta 1 )权重随着k增大而指数级衰减越久远的数据影响越小。这就是“指数移动”的含义。2.2 核心参数衰减率β与半衰期**衰减率 ( \beta ) **这是你需要理解并可能调整的最重要参数。( \beta ) 越接近1例如0.999EMA更新越“缓慢”对当前新参数 ( \theta_t ) 的响应越不敏感平滑效果越强记忆的历史信息越长。反之( \beta ) 越接近0例如0.9EMA更新越“激进”更紧跟当前参数平滑效果弱更像近期数据的简单平均。在PyTorch或TensorFlow的实现中你通常会看到另一个参数decay。这里需要小心不同库的定义可能不同。在PyTorch常见的实现如YOLOv5的EMA类中decay通常直接就是 ( \beta )。但有些地方decay可能指 ( (1 - \beta) )即当前参数的权重。所以一定要看代码的具体计算公式。半衰期这是一个更直观理解 ( \beta ) 的方式。半衰期指的是某个历史参数的权重衰减到初始权重一半所需要的步数。我们可以通过公式 ( \beta^{\text{half_life}} 0.5 ) 来估算。例如若 ( \beta 0.999 )则半衰期 ( \approx \log_{0.999}(0.5) \approx 693 ) 步。这意味着大约693个训练步batch后一个参数的“影响力”减半。这适用于训练周期很长数万步的场景。若 ( \beta 0.99 )则半衰期 ( \approx 69 ) 步。若 ( \beta 0.9 )则半衰期 ( \approx 7 ) 步。对于典型的深度学习训练例如YOLOv5在COCO上训练300个epochbatch size较大默认的 ( \beta0.9999 ) 或 ( 0.999 ) 是常见选择这能让EMA模型平滑掉数百甚至上千个batch内的波动。注意在训练初期模型参数变化剧烈EMA由于历史信息不足初始化为0或第一个参数其值会有一个“热身”阶段偏离真实移动平均。因此有些实现会引入一个偏差校正Bias Correction项尤其是在训练早期让EMA的估计更准确。不过在深度学习EMA的常见应用中因为训练步数足够多这个初始偏差的影响很快会被稀释所以很多代码如YOLOv5的EMA并没有做严格的偏差校正。3. 在PyTorch中实现与集成EMA3.1 手动实现一个基础的EMA类理解了原理我们来看如何在PyTorch中实现它。下面是一个精简但功能完整的EMA类它模仿了YOLOv5中EMA的设计思想import torch from copy import deepcopy class ModelEMA: 指数移动平均EMA模型包装器。 维护一个模型参数的影子副本并在每个训练步骤后更新它。 def __init__(self, model, decay0.9999, updates0): 初始化EMA。 Args: model: 要进行EMA的PyTorch模型。 decay: 衰减率越接近1历史权重越大平滑越强。 updates: 初始更新步数用于预热调整。 # 创建模型的深拷贝但不复制梯度计算图 self.ema_model deepcopy(model).eval() # EMA模型初始化为当前模型并设为评估模式 self.decay decay self.updates updates # 记录更新次数可用于动态调整decay # 冻结EMA模型的所有参数不参与梯度更新 for param in self.ema_model.parameters(): param.requires_grad_(False) def update(self, model): 使用当前模型参数更新EMA模型参数。 Args: model: 当前训练中的模型。 with torch.no_grad(): # 确保更新过程不计算梯度 self.updates 1 # 计算当前步的实际衰减率可选项用于预热 d self.decay * (1 - torch.exp(torch.tensor(-self.updates / 2000.0))) if self.updates 2000 else self.decay # 更新EMA模型的所有可训练参数 for ema_param, model_param in zip(self.ema_model.parameters(), model.parameters()): # 核心EMA更新公式 ema_param.mul_(d).add_(model_param.data, alpha1 - d) # 等价于: ema_param.data d * ema_param.data (1 - d) * model_param.data # 同样更新BatchNorm的running_mean和running_var如果模型有BN层 for ema_buffer, model_buffer in zip(self.ema_model.buffers(), model.buffers()): # 注意有些实现选择不更新buffer或采用不同的更新策略。YOLOv5通常会更新。 if running_mean in model_buffer._metadata or running_var in model_buffer._metadata: ema_buffer.mul_(d).add_(model_buffer.data, alpha1 - d) def __call__(self, *args, **kwargs): 使EMA模型可调用像普通模型一样进行前向推理。 return self.ema_model(*args, **kwargs) def state_dict(self): 返回EMA模型的状态字典用于保存检查点。 return self.ema_model.state_dict() def load_state_dict(self, state_dict): 加载状态字典到EMA模型。 self.ema_model.load_state_dict(state_dict)关键点解析深拷贝与评估模式初始化时用deepcopy创建主模型的完整副本作为EMA模型的起点并立即设置为.eval()模式。这是因为EMA模型只用于推理不需要计算梯度requires_gradFalse或进行Dropout/BatchNorm的统计量更新。核心更新循环update函数是核心。它遍历EMA模型和当前模型的所有参数parameters()应用EMA公式进行更新。这里使用了原地操作mul_和add_以提高效率。Buffer的更新对于BatchNorm层的running_mean和running_var它们不是参数parameters而是缓冲区buffers。是否更新它们存在争议。更新它们意味着EMA模型拥有自己独立的BN统计量这可能更一致。YOLOv5的默认实现是更新它们的。如果你发现EMA模型性能异常可以尝试注释掉buffer更新部分。预热可选代码中d的计算包含了一个简单的预热逻辑在前2000次更新中实际衰减率从0逐渐增加到设定的decay。这有助于缓解训练初期EMA值不稳定的问题。这是一个实用技巧但不是必须的。3.2 将EMA集成到训练循环中有了EMA类将其嵌入标准的PyTorch训练循环就非常直观了。以下是一个示例片段# 初始化模型、优化器、损失函数等 model YourModel().to(device) optimizer torch.optim.Adam(model.parameters(), lr1e-3) criterion nn.CrossEntropyLoss() # 初始化EMAdecay通常设为0.999或0.9999 ema ModelEMA(model, decay0.9999) # 训练循环 for epoch in range(num_epochs): model.train() # 主模型设为训练模式 for batch_idx, (data, target) in enumerate(train_loader): data, target data.to(device), target.to(device) # 1. 前向传播与损失计算 optimizer.zero_grad() output model(data) loss criterion(output, target) # 2. 反向传播与参数更新 loss.backward() optimizer.step() # 3. 更新EMA模型在optimizer.step()之后 ema.update(model) # ... 记录日志等 # 每个epoch结束后可以用EMA模型进行验证 if (epoch 1) % validation_interval 0: model.eval() # 主模型评估 # ... 主模型验证逻辑 ema.ema_model.eval() # EMA模型评估它本身一直是eval状态这里显式调用以示清晰 with torch.no_grad(): # ... 使用 ema.ema_model 进行验证计算指标 # 通常你会发现 ema.ema_model 的验证精度更高、更稳定集成要点更新时机务必在optimizer.step()更新完主模型参数之后再调用ema.update(model)。这样才能用刚更新好的最新参数去更新EMA影子。验证与保存在验证阶段你可以同时用model和ema.ema_model进行推理对比两者的性能。保存检查点时通常建议同时保存主模型和EMA模型的状态字典torch.save({ epoch: epoch, model_state_dict: model.state_dict(), ema_state_dict: ema.state_dict(), # 保存EMA模型 optimizer_state_dict: optimizer.state_dict(), loss: loss, }, checkpoint.pth)推理部署最终用于部署或测试的往往是那个在验证集上表现更好的EMA模型。你可以直接加载ema_state_dict到一个新的模型实例中。4. EMA在YOLOv5中的实战应用与调参4.1 YOLOv5中EMA的工作流在YOLOv5的代码库中EMA的实现被封装在utils/torch_utils.py的ModelEMA类中其核心逻辑与我们上面实现的基本一致。在train.py脚本中你可以看到它是如何被集成和控制的初始化在训练开始前根据命令行参数--ema默认为True决定是否创建EMA对象。训练循环更新在每个batch的反向传播和优化器更新之后调用ema.update(model)。验证与保存在验证阶段代码会分别用主模型和EMA模型在验证集上计算mAP等指标。默认情况下保存的best.pt和last.pt其实是EMA模型的权重如果启用了EMA。这是YOLOv5的一个贴心设计因为EMA模型通常更好。模型热启动如果你从预训练权重--weights开始训练YOLOv5的EMA初始化会尝试从权重文件中加载之前保存的EMA状态如果存在从而实现训练中断后的无缝恢复。4.2 关键参数调整与经验在YOLOv5的hyp.scratch.yaml或hyp.finetune.yaml超参数文件中EMA相关的配置通常是# Hyperparameters ... ema_decay: 0.9999 # EMA decay factor ema_warmup_epochs: 3.0 # EMA warmup epochsema_decay这就是我们一直讨论的衰减率β。YOLOv5默认使用0.9999这是一个非常高的值意味着平滑力度极强EMA模型参数变化非常缓慢。这个默认值在COCO这样的大数据集、长周期训练上效果很好。何时调低如果你的数据集很小例如只有几百张图片或者训练周期很短50个epoch使用0.9999可能会导致EMA模型“跟不上”主模型的快速学习。可以尝试将其降低到0.999或0.995让EMA对新参数更敏感一些。何时调高对于非常大的数据集或需要极致稳定性的场景保持0.9999或甚至0.99999可能更合适。ema_warmup_epochs热身周期数。在训练开始的这几个epoch内EMA的衰减率会从一个较小的值如0线性增加到设定的ema_decay。这有助于避免训练初期EMA值因初始化偏差而不可靠。默认3.0个epoch对于大多数情况是足够的。实操心得监控对比一个很好的习惯是在训练时同时记录主模型和EMA模型在验证集上的指标如mAP0.5。你可以用TensorBoard或WB等工具绘制两条曲线。通常你会看到EMA模型的曲线更平滑且最终收敛的“天花板”略高于或等于主模型。如果EMA模型性能显著差于主模型可能是decay设置不当或训练周期太短。小数据集的策略对于小数据集微调我个人的经验是可以尝试关闭EMA。因为小数据下模型容易过拟合参数更新本身就不稳定EMA的强平滑可能会“平滑掉”一些重要的、针对新数据集的快速适应信号。直接使用主模型进行保存和验证有时效果反而更直接。你可以通过--ema False来关闭它。最终模型选择除非有特殊原因否则默认使用EMA模型作为最终模型。在YOLOv5中这已经是默认行为。当你加载best.pt进行推理时你加载的就是EMA权重。4.3 常见问题排查FAQ1. 训练时出现CUDA内存不足OOM怀疑是EMA导致的是的EMA会创建一个完整的模型副本。虽然这个副本的requires_gradFalse但它仍然存储在GPU内存中。如果你的模型已经很大如YOLOv5x开启EMA会使显存占用几乎翻倍。解决方案使用更大的GPU或减少batch size。考虑使用--ema False关闭EMA性能可能略有损失。一些高级技巧可以将EMA权重以半精度fp16存储但实现起来较复杂需确保前向传播时类型转换正确。2. 加载保存的EMA模型进行推理发现精度不对检查加载方式确保你加载的是EMA模型的状态字典而不是主模型的。YOLOv5保存的best.pt是一个字典通常model.state_dict()对应的是EMA权重如果训练时启用了EMA。直接用torch.load(best.pt)[model].load_state_dict(...)即可。模型模式加载后务必调用model.eval()将模型设置为评估模式。这对包含BatchNorm和Dropout的模型至关重要。Buffer状态如果训练时EMA更新了BN的buffer而加载后推理结果不对可以尝试在加载权重后用一小批数据“前向传播”一次让BN层重新计算一下当前的running stats虽然理论上加载的应该是对的。或者检查你的EMA实现中buffer更新逻辑是否与训练时一致。3. EMA模型的性能在训练后期反而开始下降这种现象偶尔会发生。可能的原因过平滑衰减率decay设置过高如0.99999在训练末期模型参数已接近收敛细微的调整可能是重要的但EMA过于缓慢无法跟上这种精细调整导致“滞后”。解决方案可以尝试在训练的最后几个epoch逐步降低decay例如从0.9999线性降到0.999让EMA在末期更贴近主模型。但这需要修改训练代码实现动态decay。4. 在多GPUDataParallel/DistributedDataParallel训练中使用EMA需要注意什么模型包装确保你的EMA类是在包装之前的模型上初始化的。即model YourModel() ema ModelEMA(model) # 先初始化EMA if torch.cuda.device_count() 1: model torch.nn.DataParallel(model) # 再包装模型进行多GPU训练在update时你需要传入包装后的模型model.module来获取其原始参数# 在update函数内部或调用时 raw_model model.module if hasattr(model, module) else model self.update(raw_model)YOLOv5的EMA类内部已经处理了这种情况它会自动检测并获取model.module。5. EMA的变体与相关技术概念5.1 Stochastic Weight Averaging (SWA) 与 EMASWA是另一个著名的模型平均技术由Izmailov等人提出。它与EMA的思想有相似之处但实现方式不同EMA在每个训练步骤后更新使用固定的指数衰减持续平滑。SWA通常在训练后期以固定周期如每几个epoch保存一次模型快照最后对这些快照的权重取算术平均。对比与选择计算开销EMA每个step都需要更新有轻微计算成本SWA只在特定点保存快照平均操作通常在训练结束后进行一次计算成本更低。内存SWA需要存储多个模型快照内存开销大EMA只维护一个影子模型。效果理论上SWA通过平均多个收敛路径上的点可能找到更平坦的最小值泛化性可能更好。EMA则提供了一种连续的平滑。实践EMA实现更简单易于集成到任何训练循环中是更“傻瓜式”的默认选择。SWA需要更精细的调度何时开始收集快照。对于许多任务两者都能带来可观的提升EMA因其简便性更受欢迎。5.2 EMA与优化器状态平均我们上面讨论的EMA是对模型参数进行平均。还有一种思路是对优化器的状态进行平均例如对Adam优化器中的一阶矩和二阶矩估计进行EMA。这相当于在优化空间进行平滑有时也能带来稳定训练的好处。不过这不如参数EMA常见和通用。5.3 “EMA注意力机制”是什么在网络搜索热词中出现了“ema注意力机制”。这通常不是指我们这里讨论的用于模型权重的EMA。在计算机视觉领域有一类注意力模块被称为“EMAEfficient Multi-scale Attention”或类似名称它借鉴了移动平均的思想来高效地建模空间或通道间的长程依赖。例如在一些轻量级网络设计中会使用一维的全局移动平均池化来代替复杂的二维自注意力以降低计算量。这与本文所述的训练技巧层面的“指数移动平均”是不同的概念属于网络结构设计范畴。6. 超越YOLOv5EMA在其他深度学习任务中的实践EMA并非目标检测的专属技术它是一种通用的训练技巧适用于几乎所有使用梯度下降的深度学习模型。图像分类在训练ResNet、EfficientNet等分类网络时使用EMA同样能稳定训练提升最终测试精度。你可以在任何分类项目的训练循环中加入我们上面实现的ModelEMA类。语义分割分割网络通常也受益于EMA。训练过程的波动会影响分割边界的精细程度EMA平滑后的模型往往能产生更一致、噪声更少的预测图。自然语言处理在训练Transformer、LSTM等模型进行机器翻译、文本生成时EMA同样有效。尤其是在训练大型语言模型时参数的平滑对于生成质量的稳定性很有帮助。生成对抗网络在GAN的训练中生成器G和判别器D的对抗性训练非常不稳定。对G和D都使用EMA有时特指对G使用称为“EMA Generator”是稳定GAN训练、获得更高质量生成样本的常用trick。许多先进的GAN实现如StyleGAN2都默认使用了EMA。通用集成建议无论什么任务你都可以遵循以下模式将ModelEMA类作为通用工具放入你的项目工具包。在训练循环的optimizer.step()后立即调用ema.update(model)。在验证/测试时同时评估主模型和EMA模型选择性能更好的一个作为最终模型。将EMA模型的状态字典与主模型的一起保存。最后关于EMA衰减率的选择没有一个放之四海而皆准的值。0.999是一个很好的起点。你可以将其视为一个需要轻微调整的超参数。对于新任务可以尝试0.999, 0.9995, 0.9999这几个值并在一个小的验证集上观察哪个值带来的性能提升最稳定。记住EMA的目的是稳定和提升泛化性能而不是引入另一个需要大力调参的复杂组件。在大多数情况下采用框架的默认值如YOLOv5的0.9999并专注于数据、模型结构和基础超参数的调优会是更高效的策略。
返回列表