Метод eval
Метод eval класса Module переключает модель
в режим оценки (evaluation mode). В этом режиме
отключаются слои Dropout и BatchNorm
ведут себя иначе - используют накопленную
статистику вместо текущего батча. Метод не
принимает параметров и изменяет атрибут
training модели на False.
Синтаксис
model.eval()
Пример
Давайте создадим простую модель с Dropout и посмотрим, как меняется её поведение в режиме оценки:
import torch
import torch.nn as nn
class MyModel(nn.Module):
def __init__(self):
super().__init__()
self.dropout = nn.Dropout(0.5)
self.fc = nn.Linear(10, 10)
def forward(self, x):
x = self.dropout(x)
x = self.fc(x)
return x
model = MyModel()
model.eval()
print(model.training)
Результат выполнения кода:
False
Пример
Теперь сравним выходы Dropout в режиме обучения и оценки:
import torch
import torch.nn as nn
torch.manual_seed(0)
dropout = nn.Dropout(0.5)
x = torch.ones(1, 5)
# Режим обучения
dropout.train()
res_train = dropout(x)
print(res_train)
# Режим оценки
dropout.eval()
res_eval = dropout(x)
print(res_eval)
Результат выполнения кода:
tensor([[2., 0., 2., 2., 0.]])
tensor([[1., 1., 1., 1., 1.]])
Пример
Метод eval часто используют вместе
с контекстным менеджером no_grad
для отключения вычисления градиентов
во время инференса:
import torch
import torch.nn as nn
class SimpleModel(nn.Module):
def __init__(self):
super().__init__()
self.bn = nn.BatchNorm1d(5)
def forward(self, x):
x = self.bn(x)
return x
model = SimpleModel()
model.train()
# Обучение...
# Инференс
model.eval()
with torch.no_grad():
t = torch.randn(1, 5)
res = model(t)
print(res)
Результат выполнения кода:
tensor([[-1.4873, 0.6254, -0.3918, 1.0463, 0.2073]])