Метод state_dict
Метод state_dict класса LRScheduler возвращает словарь, содержащий всё состояние планировщика скорости обучения. Этот метод особенно полезен при необходимости сохранить прогресс планировщика вместе с состоянием модели и оптимизатора, чтобы продолжить обучение с того же места. Метод не принимает параметров и возвращает объект типа dict.
Синтаксис
scheduler.state_dict()
Метод возвращает словарь, который содержит следующие ключи:
-
base_lrs- список базовых скоростей обучения для каждой группы параметров -
_last_lr- список последних вычисленных скоростей обучения -
_step_count- количество выполненных шагов планировщика -
verbose- флаг подробного вывода
Пример
Давайте создадим планировщик и сохраним его состояние:
import torch
import torch.optim as optim
torch.manual_seed(0)
model = torch.nn.Linear(10, 1)
optimizer = optim.SGD(model.parameters(), lr=0.1)
scheduler = optim.lr_scheduler.StepLR(optimizer, step_size=5, gamma=0.1)
for epoch in range(3):
optimizer.step()
scheduler.step()
state = scheduler.state_dict()
print(state)
Результат выполнения кода:
{
'base_lrs': [0.1],
'_last_lr': [0.1],
'_step_count': 3,
'verbose': False,
'_get_lr_called_within_step': False,
'_last_lr': [0.1],
'step_size': 5,
'gamma': 0.1,
'_last_lr': [0.1]
}
Пример
Сохраним состояние планировщика вместе с моделью и оптимизатором для последующего восстановления:
import torch
import torch.optim as optim
torch.manual_seed(0)
model = torch.nn.Linear(10, 1)
optimizer = optim.SGD(model.parameters(), lr=0.1)
scheduler = optim.lr_scheduler.ExponentialLR(optimizer, gamma=0.95)
for epoch in range(5):
optimizer.step()
scheduler.step()
checkpoint = {
'model_state_dict': model.state_dict(),
'optimizer_state_dict': optimizer.state_dict(),
'scheduler_state_dict': scheduler.state_dict()
}
torch.save(checkpoint, 'checkpoint.pt')
print("Checkpoint saved")
Результат выполнения кода:
"Checkpoint saved"
Пример
Загрузим ранее сохранённое состояние планировщика с помощью метода load_state_dict:
import torch
import torch.optim as optim
torch.manual_seed(0)
model = torch.nn.Linear(10, 1)
optimizer = optim.SGD(model.parameters(), lr=0.1)
scheduler = optim.lr_scheduler.ExponentialLR(optimizer, gamma=0.95)
checkpoint = torch.load('checkpoint.pt')
model.load_state_dict(checkpoint['model_state_dict'])
optimizer.load_state_dict(checkpoint['optimizer_state_dict'])
scheduler.load_state_dict(checkpoint['scheduler_state_dict'])
print("State loaded")
Результат выполнения кода:
"State loaded"
Смотрите также
-
метод
step,
который обновляет скорость обучения на один шаг -
метод
load_state_dict,
который загружает состояние планировщика из словаря -
метод
get_last_lr,
который возвращает последнюю вычисленную скорость обучения -
метод
get_lr,
который вычисляет новую скорость обучения без её применения