欢迎光临
我们一直在努力

PyTorch 从零理解 Dropout、Early Stopping 与 Weight Decay

正则化到底怎么防过拟合?从零理解 Dropout、Early Stopping 与 Weight Decay

副标题:手把手带你建立完整的"防过拟合方法地图"——Weight Decay 限制权重、Dropout 随机关神经元、Early Stopping 及时刹车、数据增强多见世面

你辛辛苦苦训了一个模型,训练集准确率飙到 99%,兴奋地拿测试集一跑——只有 70%。模型把训练数据的每一道"错题"都背下来了,可一到新样本就露馅。这就是深度学习里最磨人的老毛病:过拟合(Overfitting)。

别急着换模型、加层数。真正该问的是:我怎么让模型"学会"而不是"背会"?答案就是正则化(Regularization)——一组专门对抗过拟合的工具。

这篇我们把最常用的几种正则化方法讲透,配上数学直觉和"学生备考"的生活化类比,让你以后面对任何任务,都知道该往工具箱里掏哪一件。


一、先用一个例子吃透"过拟合"

想象两个学生备考:

  • 学生 A:把练习册每道题的答案都死记硬背下来了。练习册上满分,换一道新题就抓瞎——这对应过拟合:模型在训练集上表现极好,但泛化(generalization)能力差。
  • 学生 B:理解了底层公式和解题思路。练习册上 90 分,遇到新题也能举一反三——这对应泛化好:模型学到了真正通用的规律。

正则化的目的,就是强行把你的模型从"学生 A"改造成"学生 B"。

怎么改?核心思想是:训练时不再只追求"训练误差最小",而是同时加上一个"模型别太复杂"的约束。写成统一公式就是:

总损失

=

原始损失

(

数据

)

+

λ

正则化项

(

模型复杂度

)

\\text{总损失} = \\text{原始损失}(数据) + \\lambda \\cdot \\text{正则化项}(模型复杂度)

总损失=原始损失(数据)+λ正则化项(模型复杂度)

  • 第一项逼模型拟合数据;
  • 第二项(正则化项)逼模型保持简单;
  • λ

    \\lambda

    λ 是平衡两者的"裁判"——

    λ

    \\lambda

    λ 越大,对复杂度的惩罚越重,模型越倾向简单。

💡 泛化误差的直觉分解:机器学习里有个经典结论,模型的泛化误差 ≈ 偏差² + 方差 + 噪声。过拟合对应的是方差太大——模型对训练样本的微小扰动极度敏感,换个采样就乱套。所有正则化手段,本质上都是在"适度抬高偏差、压低方差",找到一个泛化误差最小的平衡点。这就是为什么正则化会牺牲一点点训练精度,却换来测试精度的提升。


二、方法 1:Weight Decay(L2 正则)——别让权重长得太"野"

核心思想

权重参数如果长得太大,意味着模型对某些特征过度敏感(甚至记住了训练集里的噪声)。Weight Decay 的做法很简单粗暴:惩罚大的权重,逼它们保持小巧。

数学原理

L2 正则把"所有权重的平方和"作为惩罚项加进损失:

L

t

o

t

a

l

=

L

o

r

i

g

i

n

a

l

+

λ

2

i

w

i

2

L_{total} = L_{original} + \\frac{\\lambda}{2} \\sum_i w_i^2

Ltotal=Loriginal+2λiwi2

对权重

w

w

w 求梯度时,这一项会贡献一个

λ

w

\\lambda w

λw 的额外梯度,于是每次更新都相当于:

w

w

η

(

L

o

r

i

g

i

n

a

l

w

+

λ

w

)

=

(

1

η

λ

)

w

η

L

o

r

i

g

i

n

a

l

w

w \\leftarrow w – \\eta \\left( \\frac{\\partial L_{original}}{\\partial w} + \\lambda w \\right) = (1 – \\eta\\lambda)\\, w – \\eta \\frac{\\partial L_{original}}{\\partial w}

wwη(wLoriginal+λw)=(1ηλ)wηwLoriginal

注意那个

(

1

η

λ

)

(1 – \\eta\\lambda)

(1ηλ) 系数——它让权重每步都按比例 shrink 一点点,所以 L2 正则也叫"权重衰减(weight decay)"。这就是名字的由来。

PyTorch 代码

# AdamW 里直接设 weight_decay,最推荐
optimizer = torch.optim.AdamW(model.parameters(), lr=0.001, weight_decay=0.01)

参数常用值作用
weight_decay 0.01 约束强度

λ

\\lambda

λ,越大权重越小

🕳️ 易混淆辨析:L1 vs L2 正则

  • L2:惩罚

    w

    i

    2

    \\sum w_i^2

    wi2,梯度是

    λ

    w

    \\lambda w

    λw,让权重趋于小但非零(平滑衰减)。

  • L1:惩罚

    w

    i

    \\sum |w_i|

    wi,梯度是

    λ

    sign

    (

    w

    )

    \\lambda \\cdot \\text{sign}(w)

    λsign(w),会产生稀疏解(很多权重直接变成 0,相当于自动做特征选择)。 深度学习几乎默认用 L2(Weight Decay),因为它不会把权重清零、训练更稳定。

还有个关键细节:标准 SGD 里"L2 正则"和"weight decay"完全等价;但 Adam 优化器里两者不等价——直接在梯度里加 L2 会和 Adam 的动量/二阶矩耦合,效果变差。所以 PyTorch 的 AdamW 把 weight decay 解耦出来单独处理(这正呼应了第 54 篇讲 AdamW 时的结论),比 Adam + L2 稳定得多,是现代训练的首选。


三、方法 2:Dropout——考试前随机遮住部分笔记

核心思想(生活化类比)

训练团队时,如果某个人永远在,其他人就容易"甩锅"给他。Dropout 的做法是:每次训练随机让一部分神经元"请假",逼剩下的神经元也必须学会独立完成任务——不能过度依赖任何单一个体。

这就像你备考时随机遮住笔记的一部分去自测:遮住的地方你被迫用别的知识推导,于是整本知识你都真懂了,而不是只记得某一页。

PyTorch 代码

import torch.nn as nn

model = nn.Sequential(
nn.Linear(100, 256),
nn.ReLU(),
nn.Dropout(p=0.5), # 训练时每个神经元有 50% 概率被"关掉"
nn.Linear(256, 10)
)

参数常用值作用
p 0.1 ~ 0.5 神经元被丢弃的概率

Dropout 的数学直觉:为什么推理期要"放大补偿"?

这是 Dropout 最容易讲不清、也最值得懂的一点。

训练期:每个神经元以概率

p

p

p 被置零,以概率

1

p

1-p

1p 被保留。被保留的神经元,PyTorch 会顺手把它乘以

1

1

p

\\frac{1}{1-p}

1p1 放大。为什么?

考虑一个神经元在训练期的期望输出:

E

[

训练输出

]

=

(

1

p

)

(

原值

×

1

1

p

)

+

p

0

=

原值

\\mathbb{E}[\\text{训练输出}] = (1-p) \\cdot (\\text{原值} \\times \\tfrac{1}{1-p}) + p \\cdot 0 = \\text{原值}

E[训练输出]=(1p)(原值×1p1)+p0=原值

也就是说,乘上

1

1

p

\\frac{1}{1-p}

1p1 这步"放大",保证了训练期的期望输出 ≈ 不丢时的输出。这样训练和推理的尺度才一致。

推理期(eval):不再随机丢弃,所有神经元都工作。此时没有"丢弃带来的期望缩水",所以不需要再放大补偿——直接拿全网络的结果就行。这也是为什么 model.eval() 后 Dropout 自动关闭:它等价于"把训练期被放大的、随机的多个子网络取平均",用一个完整网络近似集成(ensemble)的效果。

# 训练时:Dropout 生效,每次前向结果都不同
model.train()
out1 = model(x)
out2 = model(x)
print(torch.equal(out1, out2)) # False,随机性导致输出不同

# 评估时:Dropout 关闭,确定性输出
model.eval()
out1 = model(x)
out2 = model(x)
print(torch.equal(out1, out2)) # True,结果稳定可复现

⚠️ 忘记 model.eval() 会让评估结果抖动甚至虚高/虚低——Dropout 还在随机丢神经元,你每次跑测试集得到的指标都不一致,根本没法比较。这是一个高频陷阱。

交互组件 1

🎮 👉 点击在线体验此交互组件

一张"Dropout 训练 vs 推理"对比图


四、方法 3:Early Stopping——见好就收,别练过头

核心思想

训练不是越久越好。看下面这条规律曲线:

训练集 loss:一直下降(模型越来越"记得"训练数据)
验证集 loss:先下降 → 达到最低点 → 开始上升(过拟合开始了)

验证集 loss 拐头向上的那一刻,就是过拟合的"起跑枪声"。Early Stopping 的逻辑就是:在枪响前后及时刹车,保存那时候最好的模型。

为什么有效?回到第一节的泛化误差分解——随着训练步数增加,模型偏差持续下降但方差持续上升;验证 loss 的最低点,恰好是两者平衡、泛化误差最小的位置。Early Stopping 相当于一个"免费的、内置的正则化器",你甚至不用调

λ

\\lambda

λ

PyTorch 代码实现

best_val_loss = float('inf')
patience = 5 # 容忍多少个 epoch 验证集不改善
no_improve = 0

for epoch in range(100):
train_loss = train_one_epoch(...) # 你的训练步骤
val_loss = evaluate(...) # 在验证集上算 loss

if val_loss < best_val_loss:
best_val_loss = val_loss
no_improve = 0
# ① 只保存"验证集最好的"那个checkpoint
torch.save(model.state_dict(), 'best_model.pth')
else:
no_improve += 1
if no_improve >= patience: # ② 连续 patience 个 epoch 没进步就停
print(f"Early stopping at epoch {epoch+1}")
break

# ③ 最终用的是"最佳时期的模型",不是最后一个 epoch 的
model.load_state_dict(torch.load('best_model.pth'))

参数常用值作用
patience 5 ~ 10 容忍几个 epoch 不改善再停

💡 关键细节:Early Stopping 保存的是 best_val_loss 对应的模型,而不是训练循环里最后一个 epoch 的模型。很多新手训完直接拿最后的权重用,结果用的是已经过拟合的那版——务必 load_state_dict 回最佳 checkpoint。

配图 2


五、方法 4:数据增强——让模型多见世面

核心思想

过拟合的一大根源是训练样本太少、太单一。数据增强的思路是:人为制造更多"略有变化"的训练样本,让模型见更多世面,从而学到本质规律而非表面细节。

图像数据增强

from torchvision import transforms

transform = transforms.Compose([
transforms.RandomHorizontalFlip(), # 随机水平翻转
transforms.RandomRotation(10), # 随机旋转 ±10 度
transforms.ColorJitter(brightness=0.2), # 亮度抖动
transforms.ToTensor(),
])

一只猫翻转、旋转后还是猫——模型被迫去学"猫的本质特征"(耳朵、胡须、轮廓),而不是记住"某张照片里猫在第几行第几列"。

文本数据增强

  • 同义词替换(把"高兴"换成"开心")
  • 随机删除停用词
  • 回译(中→英→中,借助翻译制造 paraphrasing)

⚠️ 文本增强比图像增强效果弱且更易引入噪声,要谨慎使用,别把语义改坏了。


六、方法速查表(收藏)

方法核心思想一句话常用场景
Weight Decay (L2) 限制权重太大 别太激进 深度学习必用
Dropout 随机关闭神经元 别太依赖 全连接层、NLP、CV
Early Stopping 验证集不升就停 别练过头 几乎所有任务
数据增强 制造变化样本 多见世面 图像为主
BatchNorm 稳定训练(附带轻微正则) 加点随机 CNN、深度网络
Label Smoothing 标签别写太绝对 别太自信 分类、Transformer

🕳️ 辨析:Dropout 与 BatchNorm 在推理期的不同处理。两者训练/推理行为都不同,但原因截然不同:

  • Dropout:训练期随机丢弃 + 保留神经元放大

    1

    1

    p

    \\frac{1}{1-p}

    1p1;推理期关闭丢弃、不再放大(因为期望已对齐)。

  • BatchNorm:训练期用当前 batch 的均值/方差做归一化,同时用滑动平均更新全局统计量;推理期不再用 batch 统计,而是用训练期累积的全局 running_mean / running_var。 二者都靠 model.train() / model.eval() 切换,但一个是为了"随机性",一个是为了"统计量从 batch 级切到全局级"。忘了 eval(),BN 在推理时用错统计量,Dropout 在推理时随机抖动,结果都会失真。

七、实战推荐组合(入门就照这个来)

刚开始训练,建议把下面三件套一起用上——它们正交互补,叠加收益最大:

# 1. Weight Decay(写在优化器里)
optimizer = torch.optim.AdamW(model.parameters(), lr=0.001, weight_decay=0.01)

# 2. Dropout(写在模型隐藏层之间,别放输出层后)
model = nn.Sequential(
nn.Linear(100, 256),
nn.ReLU(),
nn.Dropout(0.3), # 0.3 比 0.5 温和,新手友好
nn.Linear(256, 10)
)

# 3. Early Stopping(写在训练循环里,patience=5,保存最佳模型)

# 4. 数据增强(仅图像任务):RandomFlip / RandomRotation 等


八、三个高频错误(收藏备用)

错误 1:Dropout 加在输出层之后

# ❌ 输出层后再 Dropout,预测不稳定甚至损坏分布
model = nn.Sequential(
nn.Linear(100, 256),
nn.ReLU(),
nn.Linear(256, 10), # 输出层
nn.Dropout(0.5) # ❌ 别加在这里
)

Dropout 要加在隐藏层之间,让隐藏表示"稀疏化、去依赖";加在输出层后只会破坏最终 logits 的分布。

错误 2:Dropout 的 p 设太大

nn.Dropout(p=0.8) # 80% 都丢,信息损失过大,可能欠拟合

通常 0.1–0.5,0.3 是温和稳妥的选择。p 越大越"正则",但也越容易学不动。

错误 3:评估时忘 model.eval()

见第三节,这是 Dropout + BN 共用的致命疏漏,务必在测试/推理前切到 eval 模式并加 torch.no_grad()。


核心要点小结

  • 过拟合 = 训练好、测试差,本质是方差太大;正则化通过"抬偏差、降方差"提升泛化。
  • Weight Decay(L2):惩罚

    w

    2

    \\sum w^2

    w2,权重每步按比例衰减;用 AdamW(weight_decay=0.01) 解耦实现,比 Adam+L2 稳。L1 产生稀疏解,L2 平滑衰减。

  • Dropout:训练期随机关神经元 + 保留者放大

    1

    1

    p

    \\frac{1}{1-p}

    1p1;推理期关闭丢弃、model.eval() 自动处理,等价于对多个子网络取期望(近似集成)。

  • Early Stopping:验证 loss 拐头就停,保存最佳时期的 checkpoint(不是最后一个 epoch),是免费的正则化器。
  • 数据增强:图像任务必用,文本任务慎用;核心是"制造变化、学到本质"。
  • Dropout vs BN 推理差异:前者关随机性,后者切全局统计量,但都靠 train()/eval() 切换。
  • 入门组合:Weight Decay + Dropout(0.3) + Early Stopping,三件套叠加最稳。

  • 动手思考题

  • Dropout 在训练时丢弃 50% 神经元,评估时为什么不丢弃?那个

    1

    1

    p

    \\frac{1}{1-p}

    1p1 的放大补偿,是训练期做还是推理期做?为什么?

  • 下面配置如果要加正则化防过拟合,你应该加什么、加在哪?

    model = nn.Sequential(nn.Linear(100, 512), nn.ReLU(), nn.Linear(512, 10))
    optimizer = torch.optim.Adam(model.parameters(), lr=0.001)

  • 训练集 loss 0.1、准确率 99%,验证集 loss 0.8、准确率 75%。这是什么问题?请按优先级列出至少四种解决方案。

  • BatchNorm 和 Dropout 在 model.eval() 下行为有何不同?试着用你自己的话讲清"一个关随机、一个切统计量"。欢迎在评论区贴出你的理解,我们一起纠偏 💬

  • 下一篇,我们把前面所有训练知识串起来——怎么根据 loss 和 accuracy 曲线诊断训练问题(是欠拟合、过拟合还是别的坑)。


    📚 关于本系列

    本文是 「AI 学习路线 · 阶段四:PyTorch 深度学习基础」 系列中的一篇。所有文章在我的个人博客上都有 可交互动画 + 完整学习路线 版本,建议配合食用 👇

    🔗 在博客上阅读本文原版(含可交互组件、公式动画) 👉 正则化到底怎么防过拟合?从零理解 Dropout、Early Stopping 与 Weight Decay

    🗺️ 查看完整 AI 学习路线(从 0 到进阶,持续更新) 👉 bestsdz.xyz

    觉得有帮助的话,欢迎去博客点个收藏 ⭐,你的支持是我更新的最大动力!

    赞(0)
    未经允许不得转载:171主机测评 » PyTorch 从零理解 Dropout、Early Stopping 与 Weight Decay
    分享到: 更多 (0)

    评论 抢沙发

    • 昵称 (必填)
    • 邮箱 (必填)
    • 网址