训练流程中的配置优化器的超参数 - 权重衰减
·
训练流程中的配置优化器的超参数 - 权重衰减
flyfish
权重衰减(Weight Decay)是正则化手段,作用是约束模型权重的绝对值大小,避免权重过度拟合训练数据中的噪声。
不调用 torch.optim 的实现
权重衰减仅作用于 w,偏置 b 不参与衰减。
import torch
# 1. 生成训练数据(与原代码完全一致)
x = torch.linspace(0, 10, 100).reshape(-1, 1)
y_true = 2 * x + 3 + torch.randn_like(x) * 0.5
# 2. 初始化可学习参数
w = torch.randn(1, requires_grad=True)
b = torch.randn(1, requires_grad=True)
# 3. 超参数配置
learning_rate = 0.01
weight_decay = 0.01
epochs = 1000
# 4. 纯手动训练循环:不调用任何优化器接口
for epoch in range(epochs):
# 前向传播:计算预测值
y_pred = w * x + b
# 计算原始均方误差损失,正则项不写在损失里,直接在更新参数时体现
loss = torch.mean((y_pred - y_true) ** 2)
# 反向传播:自动计算 w.grad 和 b.grad
loss.backward()
# 手动更新参数(对应 optimizer.step())
with torch.no_grad():
# 权重 w:正常梯度更新 + 权重衰减(每一步把w往0收缩一点)
w -= learning_rate * w.grad + learning_rate * weight_decay * w
# 偏置 b:只做正常梯度更新,不施加权重衰减
b -= learning_rate * b.grad
# 手动清零梯度(对应 optimizer.zero_grad())
w.grad.zero_()
b.grad.zero_()
# 每200轮打印一次(与原代码格式一致)
if (epoch + 1) % 200 == 0:
print(f"轮次 {epoch+1:4d} | 损失: {loss.item():.4f} | w={w.item():.4f} | b={b.item():.4f}")
print(f"\n最终结果:w={w.item():.4f}, b={b.item():.4f}")
print("真实值:w=2, b=3")
调用 torch.optim 的实现
import torch
# 生成训练数据
x = torch.linspace(0, 10, 100).reshape(-1, 1)
y_true = 2 * x + 3 + torch.randn_like(x) * 0.5
# 初始化参数
w = torch.randn(1, requires_grad=True)
b = torch.randn(1, requires_grad=True)
learning_rate = 0.01
weight_decay = 0.01
epochs = 1000 # 训练轮数
# 分组优化器:w加衰减,b不加
optimizer = torch.optim.SGD([
{"params": [w], "weight_decay": weight_decay},
{"params": [b], "weight_decay": 0.0}
], lr=learning_rate)
# 训练循环
for epoch in range(epochs):
y_pred = w * x + b
loss = torch.mean((y_pred - y_true) ** 2)
optimizer.zero_grad()
loss.backward()
optimizer.step()
# 每200轮打印一次
if (epoch + 1) % 200 == 0:
print(f"轮次 {epoch+1:4d} | 损失: {loss.item():.4f} | w={w.item():.4f} | b={b.item():.4f}")
print(f"\n最终结果:w={w.item():.4f}, b={b.item():.4f}")
print("真实值:w=2, b=3")
对基础SGD优化器而言,权重衰减在数学上等价于在原始损失中加入L2正则项:
L总=L原始损失+λ2∥w∥2L_{总} = L_{原始损失} + \frac{\lambda}{2} \Vert w \Vert^2L总=L原始损失+2λ∥w∥2
其中λ\lambdaλ为权重衰减系数,数值越大,对权重的约束越强。
反向传播后,权重的梯度会多出一项 λw\lambda wλw,最终参数更新公式变为:
w=w−η⋅gw−ηλ⋅w=w⋅(1−ηλ)−η⋅gww = w - \eta \cdot g_w - \eta \lambda \cdot w = w \cdot (1-\eta\lambda) - \eta \cdot g_ww=w−η⋅gw−ηλ⋅w=w⋅(1−ηλ)−η⋅gw
直观来看,每一步更新前,权重都会先按比例向0收缩一点,权重衰减因此得名。
偏置项bbb通常不参与权重衰减,因其数值大小对过拟合影响极小,做法是仅对权重做衰减。
更多推荐

所有评论(0)