Метод 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,
который содержит список групп параметров оптимизатора