0%

EMA笔记

未加偏差修正的 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)

# 100 个样本,一个样本一个特征
X = torch.randn(100, 1)

# 真实标签
y = 2 * X + 3 + 0.1 * torch.randn(100, 1)

# 线性回归模型,输入 1 维,输出 1 维
model = nn.Linear(1, 1)

optimizer = optim.SGD(model.parameters(), lr=0.1)

criterion = nn.MSELoss()

print(list(model.parameters())) # 注意 model.parameters() 是生成器,要list()

# EMA decay系数
beta = 0.9

# EMA 参数字典
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):
# forward
y_pred = model(X)
# loss
loss = criterion(y_pred, y)
# backward
optimizer.zero_grad()
loss.backward()
optimizer.step()

# update EMA 参数
for name, param in model.named_parameters():
ema_params[name] = beta * ema_params[name] + (1 - beta) * param.data

# 偏差修正 bias correction
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())

# 这里是我又忘记.view()怎么用了,做了个测试
# y = torch.tensor([[[1,2],[3,4]],
# [[5,6],[7,8]]])
# print(y.view(2, 1, 1, 2, 1, 2))
# 从外往里数

# 画图
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()

结果是

1


但要学会调库,可以用 PyTorch 的 AveragedModel

导入和初始化:

1
2
3
4
5
from torch.optim.swa_utils import AveragedModel

# 用 avg_fn 定义 EMA 更新公式
ema_model = AveragedModel(model,
avg_fn=lambda avg_p, model_p, num_models: 0.9*avg_p + 0.1*model_p)

解释:

  1. model:原始训练模型
  2. 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
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()

# 将 EMA 参数应用到模型
model.load_state_dict(ema_model.state_dict())

# 可以恢复原模型参数
# model.load_state_dict(original_params)

解释:

  1. ema_model.state_dict() 是 EMA 参数
  2. 可以临时替换训练模型参数用于验证或部署
  3. 用完再恢复原参数,继续训练