PyTorch中model.eval()使用與作用小結(jié)
一、??model.train()與model.eval()是什么?
使用 PyTorch 進行深度學(xué)習(xí)訓(xùn)練時,我們經(jīng)常會看到如下的代碼片段:
model.train() # 訓(xùn)練階段... model.eval() # 驗證或測試階段...
很多初學(xué)者第一次看到時都會問:
“為什么要在測試前加一句 model.eval()?
不加行不行?到底起了什么作用?”
eval,英文意即為評估

在 PyTorch 中,每個神經(jīng)網(wǎng)絡(luò)模型都是一個 nn.Module 的子類。
而 nn.Module 中有兩個非常重要的模式:
| 模式 | 含義 | 常用于 |
|---|---|---|
| model.train() | 開啟訓(xùn)練模式(默認) | 模型訓(xùn)練階段 |
| model.eval() | 開啟評估模式 | 驗證、測試階段 |
?? 它們的區(qū)別不在于是否計算梯度,
而在于模型內(nèi)部某些層(如 Dropout、BatchNorm)的行為發(fā)生變化。
二、為什么需要model.eval()
神經(jīng)網(wǎng)絡(luò)中有些層在“訓(xùn)練”和“推理”階段需要不同的行為,例如:
Dropout 層
- 在訓(xùn)練時,會隨機“丟棄”一部分神經(jīng)元(防止過擬合);
- 在測試時,則應(yīng)該關(guān)閉 Dropout,讓所有神經(jīng)元都參與計算。
如果你不調(diào)用 model.eval(),
那在測試階段 Dropout 仍然會隨機丟棄神經(jīng)元,導(dǎo)致結(jié)果不穩(wěn)定、性能下降。
Batch Normalization 層(BN層)
- 在訓(xùn)練時,BatchNorm 會根據(jù)當前 mini-batch 的均值和方差進行標準化;
- 在測試時,應(yīng)該使用在訓(xùn)練中統(tǒng)計到的“全局均值和方差”來規(guī)范化。
如果不切換到 eval 模式,
BN 層會繼續(xù)更新統(tǒng)計信息,導(dǎo)致推理結(jié)果偏差甚至錯誤。
? 結(jié)論:
model.eval() 的核心作用是讓模型中某些層(Dropout、BatchNorm)進入“推理模式”。
三、model.eval()與torch.no_grad()的區(qū)別
這兩個經(jīng)常一起出現(xiàn),很多人容易混淆:
| 功能 | 是否影響 Dropout/BN | 是否停止計算梯度 | 使用場景 |
|---|---|---|---|
| model.eval() | ? 是 | ? 否 | 切換模型狀態(tài)(推理模式) |
| torch.no_grad() | ? 否 | ? 是 | 禁止梯度計算,加快推理速度、節(jié)省顯存 |
因此,推理時我們通常會這樣寫??:
model.eval() # 切換為推理模式
with torch.no_grad(): # 不計算梯度
outputs = model(inputs)
四、完整示例:對比train()和eval()
讓我們用一個小例子直觀看看區(qū)別 ??
import torch
import torch.nn as nn
# 一個簡單的網(wǎng)絡(luò),包含 Dropout
class SimpleNet(nn.Module):
def __init__(self):
super().__init__()
self.fc = nn.Linear(4, 4)
self.dropout = nn.Dropout(p=0.5)
def forward(self, x):
return self.dropout(self.fc(x))
# 創(chuàng)建模型和輸入
x = torch.ones(4)
model = SimpleNet()
# 訓(xùn)練模式
model.train()
print("Train Mode Output:")
for _ in range(3):
print(model(x))
# 推理模式
model.eval()
print("\nEval Mode Output:")
for _ in range(3):
print(model(x))
輸出對比
Train Mode Output: tensor([-0.0000, -1.4387, 0.7793, 0.0000], grad_fn=<MulBackward0>) tensor([-0.0000, -1.4387, 0.0000, 0.0000], grad_fn=<MulBackward0>) tensor([-0.0000, -1.4387, 0.7793, 0.0000], grad_fn=<MulBackward0>) Eval Mode Output: tensor([-0.2442, -0.7194, 0.3897, 0.9389], grad_fn=<ViewBackward0>) tensor([-0.2442, -0.7194, 0.3897, 0.9389], grad_fn=<ViewBackward0>) tensor([-0.2442, -0.7194, 0.3897, 0.9389], grad_fn=<ViewBackward0>)
? 說明:
- 訓(xùn)練模式下 Dropout 隨機屏蔽神經(jīng)元,因此每次輸出不同;
- 推理模式下 Dropout 被關(guān)閉,輸出穩(wěn)定。
五、與model.train()的區(qū)別總結(jié)
| 比較項 | model.train() | model.eval() |
|---|---|---|
| 模型狀態(tài) | 訓(xùn)練模式 | 推理模式 |
| Dropout | 啟用隨機丟棄 | 關(guān)閉 |
| BatchNorm | 使用批次統(tǒng)計 | 使用全局統(tǒng)計 |
| 是否影響梯度 | ? 不影響 | ? 不影響 |
| 常用場景 | 模型訓(xùn)練階段 | 驗證、推理階段 |
六、完整實戰(zhàn)代碼(訓(xùn)練 + 驗證)
import torch
import torch.nn as nn
import torch.optim as optim
# 定義簡單模型
class Net(nn.Module):
def __init__(self):
super().__init__()
self.fc1 = nn.Linear(4, 10)
self.relu = nn.ReLU()
self.dropout = nn.Dropout(0.5)
self.fc2 = nn.Linear(10, 3)
def forward(self, x):
x = self.fc1(x)
x = self.relu(x)
x = self.dropout(x)
return self.fc2(x)
model = Net()
optimizer = optim.Adam(model.parameters())
criterion = nn.CrossEntropyLoss()
for epoch in range(3):
# ===== 訓(xùn)練階段 =====
model.train()
optimizer.zero_grad()
x = torch.randn(5, 4)
y = torch.randint(0, 3, (5,))
out = model(x)
loss = criterion(out, y)
loss.backward()
optimizer.step()
# ===== 驗證階段 =====
model.eval()
with torch.no_grad():
val_x = torch.randn(5, 4)
val_out = model(val_x)
val_pred = val_out.argmax(dim=1)
print(f"Epoch {epoch}: loss={loss.item():.4f}, val_pred={val_pred.tolist()}")
輸出如下:
Epoch 0: loss=1.0044, val_pred=[1, 1, 1, 2, 2] Epoch 1: loss=0.9953, val_pred=[2, 1, 2, 2, 2] Epoch 2: loss=1.2143, val_pred=[2, 2, 1, 2, 1]
? 訓(xùn)練時:
- Dropout 啟用;
- BatchNorm 統(tǒng)計更新。
? 驗證時:
- Dropout 關(guān)閉;
- BatchNorm 使用訓(xùn)練統(tǒng)計參數(shù)。
七、常見錯誤與避坑指南
| 錯誤用法 | 后果 |
|---|---|
| 在測試時忘記 model.eval() | Dropout、BN 層仍隨機,導(dǎo)致結(jié)果波動、不穩(wěn)定 |
| 在推理時忘記 torch.no_grad() | 會記錄梯度,浪費顯存、速度變慢 |
| 在訓(xùn)練時調(diào)用了 model.eval() | 模型學(xué)不動,BN 不更新統(tǒng)計信息 |
| 忘記在訓(xùn)練開始前加 model.train() | 模型仍在推理模式,訓(xùn)練效果不佳 |
八、小結(jié)
| 項目 | 說明 |
|---|---|
| 函數(shù)名 | model.eval() |
| 所屬模塊 | torch.nn.Module |
| 作用 | 切換模型到評估(推理)模式 |
| 影響層 | Dropout、BatchNorm |
| 與 no_grad 區(qū)別 | eval() 控制模式,no_grad 控制梯度 |
| 使用場景 | 驗證、測試、推理階段 |
| 常用組合 | model.eval() + with torch.no_grad(): |
到此這篇關(guān)于PyTorch中model.eval()使用與作用小結(jié)的文章就介紹到這了,更多相關(guān)PyTorch model.eval()使用內(nèi)容請搜索腳本之家以前的文章或繼續(xù)瀏覽下面的相關(guān)文章希望大家以后多多支持腳本之家!
相關(guān)文章
如何利用python多線程爬取天氣網(wǎng)站圖片并保存
最近做個天 氣方面的APP需要用到一些天氣數(shù)據(jù),所以下面這篇文章主要給大家介紹了關(guān)于如何利用python多線程爬取天氣網(wǎng)站圖片并保存的相關(guān)資料,文中通過示例代碼介紹的非常詳細,需要的朋友可以參考下2021-11-11
Python之tkinter列表框Listbox與滾動條Scrollbar解讀
這篇文章主要介紹了Python之tkinter列表框Listbox與滾動條Scrollbar解讀,具有很好的參考價值,希望對大家有所幫助。如有錯誤或未考慮完全的地方,望不吝賜教2023-05-05
Python利用pynput實現(xiàn)劃詞復(fù)制功能
這篇文章主要為大家想詳細介紹了Python如何利用pynput實現(xiàn)劃詞復(fù)制功能,文中的示例代碼講解詳細,感興趣的小伙伴可以跟隨小編一起學(xué)習(xí)一下2022-05-05
PyTorch中關(guān)于tensor.repeat()的使用
這篇文章主要介紹了PyTorch中關(guān)于tensor.repeat()的使用,具有很好的參考價值,希望對大家有所幫助。如有錯誤或未考慮完全的地方,望不吝賜教2022-11-11

