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

Метод 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,
    который содержит группы параметров оптимизатора
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить