Model EMA 在训练全程额外维护一份指数移动平均的影子权重用于交付:一行递推、不建计算图、不参与反传,把权重轨迹的高频抖动滤掉——训练走优化器,采样用影子。

影子权重:对权重轨迹做平滑

EMA 的数学内核见指数移动平均;Model EMA 把平滑对象从梯度换成模型权重。训练全程额外维护一份影子权重 ,每步做 ,不建计算图、不参与反传。mini-batch 梯度天然带噪声,原始权重在损失面上震荡前进;EMA 对最近一段轨迹加权平均,落点更接近平坦区域。生成模型(扩散、AR 采样)收益最大——推理输出对权重微小扰动敏感,EMA 权重的采样稳定性提升肉眼可见。

思想上这是 Polyak averaging 的指数折扣版:等权轨迹平均换成指数衰减权,换来一份状态、一个递推步的在线实现。滤波器视角下就是给权重轨迹接一只一阶低通:降噪与滞后都由 决定,影子权重总是”落后”训练进度约一个有效窗口。

GLM4V 训练配置

值 / 含义
decay0.9999,有效窗口 ≈ 10000 步
存储形式训练 ckpt 字典里独立的 ema_model 字段
融合实现fused_ema_adamw 把 AdamW 更新和 EMA 更新合成一个 NPU 融合算子,省一次全参数遍历
推理取用convert_glm4v_tp_new.py --use_emaema_model 字段做 TP=1 合并,采样用的就是 EMA 权重

判断 ckpt 是否带 EMA

不能看 checkpoint 目录名后缀,要 torch.load 之后查字典里有没有 ema_model 这个 key——EMA 存在 ckpt 内部,文件名不承载这个信息。

最小实现与工程代价

剥掉融合优化的裸逻辑,一个最小实现长这样:

import torch
 
@torch.no_grad()
def ema_update(shadow: torch.nn.Module, model: torch.nn.Module, decay: float = 0.9999):
    # 影子权重只做凸组合:不建图、不反传、不动梯度
    for s, p in zip(shadow.parameters(), model.parameters()):
        s.mul_(decay).add_(p.detach(), alpha=1.0 - decay)

工业实现把这一步和优化器更新融合成单个算子,全参数遍历从两次降到一次,EMA 的计算开销几乎被吃掉,只剩一份额外显存占用。

两份权重并存的设计值得记牢:训练指标看 raw ,它代表当前优化进度;交付采样用 ,它代表近期轨迹的平均水平——训练走 AdamW,交付用 EMA。和学习率预热对照着看更清楚:预热治训练初期的不稳定,EMA 治训练后期的抖动,一个作用在梯度侧,一个作用在权重侧。

相关