深度学习中模型的不同模式
🧠 深度学习中模型的模式切换(model.train() vs model.eval())讲解
在 PyTorch 中,模型通常有两种模式:
model.train():训练模式model.eval():评估模式
下面我们从机制、源码、实践、底层逻辑等角度深入剖析。
🚀 一、模型模式的本质 —— 状态切换器
model.train() 和 model.eval() 本质是对模型中 self.training 标志位的设置,并递归传播到所有子模块。
def train(self: T, mode: bool = True) -> T:
self.training = mode
for module in self.children():
module.train(mode)
return self
这就构成了训练/评估模式切换的核心机制。
🎯 二、模式切换影响的模块
并不是所有模块都对模式切换敏感,主要有以下两类:
✅ 1. Dropout(torch.nn.Dropout)
| 模式 | 行为 |
|---|---|
train() | 随机将部分神经元输出置为 0(抑制过拟合) |
eval() | 不做任何丢弃,保留所有神经元 |
📌 注意:Dropout 在评估模式下仍然会对输出进行缩放,保持期望一致。
✅ 2. BatchNorm(torch.nn.BatchNorm2d 等)
| 模式 | 行为 |
|---|---|
train() | 用当前 batch 的均值和方差进行归一化,并更新统计值 |
eval() | 使用训练阶段累计的 running mean 和 var 归一化 |
🧵 三、源码级剖析(以 Dropout 为例)
Dropout 的核心前向传播逻辑如下:
def forward(self, input):
return F.dropout(input, self.p, self.training, self.inplace)
当 model.train() 时,self.training=True,Dropout 会启用;当 model.eval() 时,Dropout 关闭。
🧪 四、训练 & 验证的实践模板
for epoch in range(num_epochs):
# 训练阶段
model.train()
for inputs, targets in train_loader:
optimizer.zero_grad()
outputs = model(inputs)
loss = criterion(outputs, targets)
loss.backward()
optimizer.step()
# 验证阶段
model.eval()
total_loss = 0
with torch.no_grad():
for inputs, targets in val_loader:
outputs = model(inputs)
loss = criterion(outputs, targets)
total_loss += loss.item()
💥 五、常见误区解析
❌ 误区 1:以为 model.eval() 会冻结参数
不会。是否更新参数由 quires.grad控制,或者通过 torch.no_grad()`。
❌ 误区 2:推理/验证时忘了 model.eval()
这会导致 Dropout 仍然随机丢弃,BatchNorm 也会用 batch 的统计量,验证结果不稳定。
❌ 误区 3:只设置顶层模块模式就够了
不行,必须递归所有子模块,model.train() 和 model.eval() 已经做了这件事。
🧬 六、model.eval() vs with torch.no_grad()
| API | 控制对象 | 需要一起使用? |
|---|---|---|
model.eval() | 控制 Dropout、BatchNorm 行为 | 是 |
with torch.no_grad() | 控制是否计算梯度(影响显存和速度) | 是 |
典型用法如下:
model.eval()
with torch.no_grad():
outputs = model(inputs)
🧱 七、与 .to(device)、多卡训练的关系
✅ .to(device) 与模式无关
model = MyModel().to("cuda")
只改变模型计算设备,不影响训练/评估行为。
✅ 多卡训练时模式仍需设置
model = torch.nn.DataParallel(model)
model.train() # 模式设置仍然有效
🧠 八、总结金句
model.train()/model.eval()改变模型的行为,不改变结构。- 模式切换主要影响 Dropout 和 BatchNorm。
- 推理/验证阶段必须使用
eval()和torch.no_grad()。 - 模式切换不影响是否计算梯度,也不控制参数冻结。
📓 可选深入(可继续展开)
-
Dropout 数学期望推导及缩放公式:
y={0,以概率 p丢弃x1−p,以概率 1−p保留 y = \begin{cases} 0, & \text{以概率 } p \text{丢弃} \\ \frac{x}{1 - p}, & \text{以概率 } 1 - p \text{保留} \end{cases} y={0,1−px,以概率 p丢弃以概率 1−p保留
保证 E[y]=x\mathbb{E}[y] = xE[y]=x -
BatchNorm 的运行均值更新方式:
μrunning=(1−momentum)⋅μrunning+momentum⋅μbatch \mu_{\text{running}} = (1 - \text{momentum}) \cdot \mu_{\text{running}} + \text{momentum} \cdot \mu_{\text{batch}} μrunning=(1−momentum)⋅μrunning+momentum⋅μbatch -
自定义模块中根据
self.training判断模式:
class MyModule(nn.Module):
def forward(self, x):
if self.training:
# 训练逻辑
else:
# 评估逻辑
更多推荐

所有评论(0)