РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
453 of 769 menu

Метод 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]])

Смотрите также

  • метод train,
    который переключает модель в режим обучения
  • атрибут training,
    который хранит текущий режим модели
  • метод forward,
    который определяет прямой проход модели
  • класс Module,
    базовый класс для всех нейросетевых модулей
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить