未加偏差修正的 EMA:
- $\theta_t$ :当前模型参数
- $\theta^{EMA}_t$ :EMA 参数(在影子模型(shadow model)里)
- $\beta \in [0,1)$ :衰减系数接近 1 (越接近 1,衰减得越慢,历史参数影响保持更久,EMA 越平滑)
- 直觉:靠前的参数权重小,靠后的参数权重大
列举来感觉一下,若$\beta = 0.9$,则第t轮里,最新值$x_t$、前一轮的$x_{t-1}$占的权重可以计算:
- 最新值 $x_t$ 占 $1-\beta = 0.1$。
- 上一轮 $x_{t-1}$ 占 $\beta \cdot (1-\beta) = 0.9 \times 0.1 = 0.09$。
权重会衰减下去:
这就是 “指数” 两个字的由来:越早的值权重越小,按指数方式衰减。
展开前几步:
这就是一个 几何级数:
- 最新值权重最大 ($1-\beta$)
- 越早的值,权重按 $\beta, \beta^2, \dots$ 递减
- 所有权重加起来 ≈ 1(证明:等比数列求和)
每往前推一步,权重就多乘一个 $\beta$,所以是指数衰减。
因为 EMA 初期从 0 开始,会偏小,需要偏差修正(bias correction):
- $\hat{\theta}_t$ :修正后的 EMA
- 第 1 轮影响最大,随着 $t$ 增加,修正效果减小
- 直观:把初期 EMA 拉回合理水平,更快接近模型参数趋势
我的代码实现(只写了一个适用于Linear的):
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80
| import torch import torch.nn as nn import torch.optim as optim
torch.manual_seed(233)
X = torch.randn(100, 1)
y = 2 * X + 3 + 0.1 * torch.randn(100, 1)
model = nn.Linear(1, 1)
optimizer = optim.SGD(model.parameters(), lr=0.1)
criterion = nn.MSELoss()
print(list(model.parameters()))
beta = 0.9
ema_params = {} for name, param in model.named_parameters(): ema_params[name] = param.data.clone()
model_weights = [] ema_weights = [] ema_corr_weights = []
for epoch in range(1, 20): y_pred = model(X) loss = criterion(y_pred, y) optimizer.zero_grad() loss.backward() optimizer.step()
for name, param in model.named_parameters(): ema_params[name] = beta * ema_params[name] + (1 - beta) * param.data
correct_ema = {} for name in ema_params: correct_ema[name] = ema_params[name] / (1 - beta ** epoch)
print(f"Epoch{epoch}:") for name, param in model.named_parameters(): print(f"参数{name}: ①model: {param.data.view(-1)} ②EMA(corrected): {correct_ema[name].view(-1)}") print("-"*50)
model_weights.append(model.weight.data.item()) ema_weights.append(ema_params['weight'].item()) ema_corr_weights.append(correct_ema['weight'].item())
import matplotlib.pyplot as plt plt.figure(figsize=(15,8)) plt.plot(model_weights, label='Model Weight') plt.plot(ema_weights, label='EMA') plt.plot(ema_corr_weights, label='EMA Corrected') plt.title('Weight Parameter vs EMA') plt.xlabel('Epoch') plt.ylabel('Weight') plt.legend() plt.grid(True) plt.show()
|
结果是

但要学会调库,可以用 PyTorch 的 AveragedModel 。
导入和初始化:
1 2 3 4 5
| from torch.optim.swa_utils import AveragedModel
ema_model = AveragedModel(model, avg_fn=lambda avg_p, model_p, num_models: 0.9*avg_p + 0.1*model_p)
|
解释:
model:原始训练模型
avg_fn:自定义函数,告诉 AveragedModel 每轮怎么更新影子参数
avg_p:影子参数(之前 EMA 的值)
model_p:当前模型参数
num_models:默认是累积模型数,但 EMA 不需要用到
- 返回值就是新的 EMA 参数
- 参数要按照顺序写
这里 0.9*avg_p + 0.1*model_p 就是 EMA 更新公式,β=0.9
训练时更新 EMA:
1 2 3 4 5 6 7 8 9 10
| for epoch in range(num_epochs): optimizer.zero_grad() y_pred = model(X) loss = criterion(y_pred, y) loss.backward() optimizer.step() ema_model.update_parameters(model)
|
解释:
update_parameters(model) 会把当前模型参数 按 avg_fn 更新到影子模型
- 注意:EMA 的更新在训练之后,影子模型不会影响训练,只是记录趋势
验证时使用 EMA 参数:
1 2 3 4 5 6 7 8
| original_params = model.state_dict()
model.load_state_dict(ema_model.state_dict())
|
解释:
ema_model.state_dict() 是 EMA 参数
- 可以临时替换训练模型参数用于验证或部署
- 用完再恢复原参数,继续训练