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

Метод state_dict

Метод state_dict оптимизатора возвращает словарь Python, содержащий всё состояние оптимизатора. Этот словарь включает параметры состояния для каждого тензора (например, моменты в алгоритме Adam), а также гиперпараметры и текущее состояние групп параметров. Метод не принимает аргументов и используется для сохранения оптимизатора при чекпоинтах модели.

Синтаксис

optimizer.state_dict()

Пример

Давайте создадим простую модель и оптимизатор, а затем сохраним его состояние в словарь:

import torch model = torch.nn.Linear(5, 2) optimizer = torch.optim.SGD(model.parameters(), lr=0.01) state = optimizer.state_dict() print(type(state))

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

<class 'dict'>

Пример

Теперь выполним несколько шагов оптимизации, чтобы состояние заполнилось данными, и посмотрим структуру словаря:

import torch torch.manual_seed(0) model = torch.nn.Linear(5, 2) optimizer = torch.optim.SGD(model.parameters(), lr=0.01) x = torch.randn(3, 5) y = torch.randn(3, 2) loss = torch.nn.functional.mse_loss(model(x), y) loss.backward() optimizer.step() state = optimizer.state_dict() print(list(state.keys()))

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

['state', 'param_groups']

Пример

Рассмотрим структуру словаря подробнее. Ключ state содержит состояния параметров, а param_groups - группы параметров и их гиперпараметры:

import torch torch.manual_seed(0) model = torch.nn.Linear(5, 2) optimizer = torch.optim.SGD(model.parameters(), lr=0.01, momentum=0.9) x = torch.randn(3, 5) y = torch.randn(3, 2) loss = torch.nn.functional.mse_loss(model(x), y) loss.backward() optimizer.step() state_dict = optimizer.state_dict() print("State dict keys:") for key in state_dict['state']: print(f" Parameter {key}: {list(state_dict['state'][key].keys())}") print("Parameter groups keys:") for group in state_dict['param_groups']: print(f" Group: {list(group.keys())}")

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

"State dict keys: Parameter 0: ['momentum_buffer'] Parameter 1: ['momentum_buffer'] Parameter groups keys: Group: ['params', 'lr', 'momentum', 'dampening', 'weight_decay', 'nesterov']"

Пример

Метод state_dict часто используется вместе с методом load_state_dict для полного восстановления состояния оптимизатора:

import torch torch.manual_seed(0) model = torch.nn.Linear(5, 2) optimizer = torch.optim.Adam(model.parameters(), lr=0.001) for i in range(3): x = torch.randn(3, 5) y = torch.randn(3, 2) loss = torch.nn.functional.mse_loss(model(x), y) loss.backward() optimizer.step() optimizer.zero_grad() state = optimizer.state_dict() optimizer2 = torch.optim.Adam(model.parameters(), lr=0.001) optimizer2.load_state_dict(state) print("State restored") print(f"Learning rate: {optimizer2.param_groups[0]['lr']}")

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

"State restored Learning rate: 0.001"

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

  • метод load_state_dict,
    который восстанавливает состояние оптимизатора из словаря
  • метод step,
    который выполняет один шаг оптимизации параметров
  • метод zero_grad,
    который обнуляет градиенты всех параметров
  • атрибут param_groups,
    который содержит список групп параметров оптимизатора
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить