Метод load_state_dict
Метод load_state_dict класса Optimizer загружает состояние оптимизатора из словаря, полученного ранее с помощью метода state_dict. Это позволяет восстанавливать состояние оптимизатора после сохранения, что особенно полезно при возобновлении обучения модели. Первым параметром метод принимает словарь состояния, вторым параметром можно передать флаг строгого соответствия.
Синтаксис
optimizer.load_state_dict(state_dict, strict=True)
Параметры метода:
-
state_dict- словарь состояния, полученный изstate_dict -
strict- булевый флаг, указывающий на необходимость строгого соответствия ключей
Пример сохранения и загрузки состояния
Давайте создадим модель, оптимизатор и сохраним их состояния, а затем загрузим оптимизатор из сохранённого состояния:
import torch
import torch.nn as nn
import torch.optim as optim
torch.manual_seed(0)
# Создаём модель
model = nn.Linear(10, 5)
optimizer = optim.SGD(model.parameters(), lr=0.01, momentum=0.9)
# Сохраняем состояния
state = {
'model_state': model.state_dict(),
'optimizer_state': optimizer.state_dict(),
}
torch.save(state, 'model.pt')
# Создаём новую модель и оптимизатор
new_model = nn.Linear(10, 5)
new_optimizer = optim.SGD(new_model.parameters(), lr=0.01, momentum=0.9)
# Загружаем состояния
checkpoint = torch.load('model.pt')
new_model.load_state_dict(checkpoint['model_state'])
new_optimizer.load_state_dict(checkpoint['optimizer_state'])
print("Optimizer state loaded successfully")
Результат выполнения кода:
"Optimizer state loaded successfully"
Пример с параметром strict
Параметр strict определяет, нужно ли проверять точное соответствие ключей словаря состояния. Рассмотрим пример использования:
import torch
import torch.nn as nn
import torch.optim as optim
torch.manual_seed(0)
model = nn.Linear(10, 5)
optimizer = optim.SGD(model.parameters(), lr=0.01)
# Сохраняем состояние
state_dict = optimizer.state_dict()
# Создаём новый оптимизатор с другими параметрами
new_model = nn.Linear(10, 5)
new_optimizer = optim.Adam(new_model.parameters(), lr=0.001)
# Пытаемся загрузить состояние с strict=False
new_optimizer.load_state_dict(state_dict, strict=False)
print("State loaded with strict=False")
Результат выполнения кода:
"State loaded with strict=False"
Пример полного цикла обучения с сохранением состояния
Создадим полный цикл обучения с сохранением состояния оптимизатора и последующей загрузкой:
import torch
import torch.nn as nn
import torch.optim as optim
torch.manual_seed(0)
# Создаём модель и оптимизатор
model = nn.Linear(10, 1)
optimizer = optim.SGD(model.parameters(), lr=0.01)
criterion = nn.MSELoss()
# Генерируем данные
x = torch.randn(100, 10)
y = torch.randn(100, 1)
# Обучаем модель
for epoch in range(5):
optimizer.zero_grad()
output = model(x)
loss = criterion(output, y)
loss.backward()
optimizer.step()
if epoch == 2:
torch.save({
'model_state_dict': model.state_dict(),
'optimizer_state_dict': optimizer.state_dict(),
}, 'checkpoint.pt')
# Восстанавливаем модель и оптимизатор из чекпоинта
new_model = nn.Linear(10, 1)
new_optimizer = optim.SGD(new_model.parameters(), lr=0.01)
checkpoint = torch.load('checkpoint.pt')
new_model.load_state_dict(checkpoint['model_state_dict'])
new_optimizer.load_state_dict(checkpoint['optimizer_state_dict'])
print("Model and optimizer restored from checkpoint")
Результат выполнения кода:
"Model and optimizer restored from checkpoint"
Смотрите также
-
метод
state_dict,
который возвращает словарь состояния оптимизатора -
метод
step,
который выполняет шаг оптимизации -
метод
zero_grad,
который обнуляет градиенты параметров модели -
атрибут
param_groups,
который содержит группы параметров оптимизатора