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