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

Метод train

Метод train класса Module переключает модель и все её подмодули в режим обучения. Этот режим влияет на слои, которые ведут себя по-разному на этапах обучения и оценки, такие как Dropout и BatchNorm. Первым параметром метод принимает булевое значение mode. По умолчанию mode равен True, что устанавливает режим обучения.

Синтаксис

model.train(mode=True)

Пример

Давайте создадим простую модель и проверим её режим работы:

import torch import torch.nn as nn class Net(nn.Module): def __init__(self): super().__init__() self.fc = nn.Linear(10, 5) self.dropout = nn.Dropout(0.5) def forward(self, x): x = self.fc(x) x = self.dropout(x) return x model = Net() print(model.training)

Результат выполнения кода:

True

Пример

Переключим модель в режим оценки с помощью метода eval, а затем вернём обратно в режим обучения с помощью train:

import torch import torch.nn as nn class Net(nn.Module): def __init__(self): super().__init__() self.fc = nn.Linear(10, 5) self.dropout = nn.Dropout(0.5) def forward(self, x): x = self.fc(x) x = self.dropout(x) return x model = Net() model.eval() print(model.training) model.train() print(model.training)

Результат выполнения кода:

False True

Пример

Метод train можно вызвать с параметром False, чтобы отключить режим обучения для всех подмодулей:

import torch import torch.nn as nn class SubNet(nn.Module): def __init__(self): super().__init__() self.fc = nn.Linear(5, 2) class Net(nn.Module): def __init__(self): super().__init__() self.sub = SubNet() self.dropout = nn.Dropout(0.5) def forward(self, x): x = self.sub(x) x = self.dropout(x) return x model = Net() model.train(False) print(model.training) print(model.sub.training)

Результат выполнения кода:

False False

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

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