Метод load_state_dict
Метод load_state_dict класса LRScheduler восстанавливает внутреннее состояние планировщика скорости обучения из словаря состояний, полученного ранее с помощью метода state_dict. Это необходимо при возобновлении обучения после сохранения состояния оптимизатора и планировщика. Метод принимает один обязательный параметр - словарь состояний, и возвращает None.
Синтаксис
scheduler.load_state_dict(state_dict)
Пример
Давайте создадим простой планировщик и сохраним его состояние, а затем восстановим его из словаря:
import torch
from torch.optim import SGD
from torch.optim.lr_scheduler import StepLR
# Create model and optimizer
model = torch.nn.Linear(5, 1)
optimizer = SGD(model.parameters(), lr=0.1)
scheduler = StepLR(optimizer, step_size=5, gamma=0.5)
# Perform some steps
for epoch in range(3):
scheduler.step()
# Save state dictionary
state = scheduler.state_dict()
print(state)
Результат выполнения кода:
{
'step_size': 5,
'gamma': 0.5,
'base_lrs': [0.1],
'last_epoch': 3,
'_step_count': 3,
'verbose': False,
'_get_lr_called_within_step': False,
'_last_lr': [0.1]
}
Пример
Теперь создадим новый планировщик и восстановим его состояние из сохранённого словаря:
import torch
from torch.optim import SGD
from torch.optim.lr_scheduler import StepLR
# Create new model and optimizer
model = torch.nn.Linear(5, 1)
optimizer = SGD(model.parameters(), lr=0.1)
scheduler_new = StepLR(optimizer, step_size=5, gamma=0.5)
# Restore state from dictionary (assuming 'state' from previous example)
state = {
'step_size': 5,
'gamma': 0.5,
'base_lrs': [0.1],
'last_epoch': 3,
'_step_count': 3,
'verbose': False,
'_get_lr_called_within_step': False,
'_last_lr': [0.1]
}
scheduler_new.load_state_dict(state)
print(scheduler_new.get_last_lr())
print(scheduler_new.last_epoch)
Результат выполнения кода:
[0.1]
3
Пример
Восстановление состояния планировщика после сохранения вместе с оптимизатором позволяет возобновить обучение с того же места:
import torch
from torch.optim import SGD
from torch.optim.lr_scheduler import StepLR
torch.manual_seed(0)
# Initial setup
model = torch.nn.Linear(5, 1)
optimizer = SGD(model.parameters(), lr=0.1)
scheduler = StepLR(optimizer, step_size=2, gamma=0.5)
# Train for 3 epochs
for epoch in range(3):
optimizer.step()
scheduler.step()
print(f"Epoch {epoch+1}, lr: {scheduler.get_last_lr()[0]}")
# Save state
scheduler_state = scheduler.state_dict()
optimizer_state = optimizer.state_dict()
# Create new model and restore states
new_model = torch.nn.Linear(5, 1)
new_optimizer = SGD(new_model.parameters(), lr=0.1)
new_scheduler = StepLR(new_optimizer, step_size=2, gamma=0.5)
new_optimizer.load_state_dict(optimizer_state)
new_scheduler.load_state_dict(scheduler_state)
# Continue training
for epoch in range(3, 5):
new_optimizer.step()
new_scheduler.step()
print(f"Epoch {epoch+1}, lr: {new_scheduler.get_last_lr()[0]}")
Результат выполнения кода:
Epoch 1, lr: 0.1
Epoch 2, lr: 0.05
Epoch 3, lr: 0.05
Epoch 4, lr: 0.025
Epoch 5, lr: 0.0125
Смотрите также
-
метод
state_dict,
который возвращает словарь состояний планировщика -
метод
step,
который обновляет скорость обучения после каждой эпохи -
метод
get_last_lr,
который возвращает последние использованные скорости обучения -
класс
LRScheduler,
базовый класс для всех планировщиков скорости обучения