Метод load_state_dict
Метод load_state_dict класса GradScaler восстанавливает состояние объекта из словаря, ранее сохранённого с помощью метода state_dict. Это позволяет продолжить обучение модели с теми же настройками масштабирования градиентов, которые были в момент сохранения. Метод принимает один обязательный параметр - словарь состояния и возвращает объект типа NamedTuple с информацией о результате загрузки.
Основное применение load_state_dict - восстановление состояний градиентного скалера при возобновлении обучения после сохранения чекпоинтов. Это обеспечивает согласованность процесса масштабирования градиентов на протяжении всего обучения.
Синтаксис
grad_scaler.load_state_dict(state_dict)
Параметры
Метод load_state_dict принимает следующие параметры:
-
⁅b⁆state_dict⁅/b⁆ - словарь состояния, полученный ранее из метода
state_dict. Должен содержать все ключи, необходимые для полного восстановления состояния скалера.
Возвращаемое значение
Метод возвращает объект типа NamedTuple с полем missing_keys (список отсутствующих ключей) и unexpected_keys (список неожиданных ключей). Это позволяет отследить успешность загрузки состояния.
Пример
Давайте создадим градиентный скалер, сохраним его состояние, а затем восстановим из сохранённого словаря:
import torch
from torch.cuda.amp import GradScaler
torch.manual_seed(0)
scaler = GradScaler()
state_dict = scaler.state_dict()
print(type(state_dict))
Результат выполнения кода:
<class 'dict'>
А теперь загрузим сохранённое состояние:
import torch
from torch.cuda.amp import GradScaler
torch.manual_seed(0)
scaler = GradScaler()
state_dict = scaler.state_dict()
new_scaler = GradScaler()
result = new_scaler.load_state_dict(state_dict)
print(result)
Результат выполнения кода:
LoadStateDictReturn(missing_keys=[], unexpected_keys=[])
Пример
Рассмотрим полный процесс сохранения и восстановления состояния скалера вместе с моделью:
import torch
from torch.cuda.amp import GradScaler, autocast
torch.manual_seed(0)
model = torch.nn.Linear(10, 1)
optimizer = torch.optim.SGD(model.parameters(), lr=0.01)
scaler = GradScaler()
for epoch in range(2):
data = torch.randn(4, 10)
target = torch.randn(4, 1)
optimizer.zero_grad()
with autocast():
output = model(data)
loss = torch.nn.functional.mse_loss(output, target)
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
checkpoint = {
'model_state_dict': model.state_dict(),
'optimizer_state_dict': optimizer.state_dict(),
'scaler_state_dict': scaler.state_dict(),
}
new_model = torch.nn.Linear(10, 1)
new_optimizer = torch.optim.SGD(new_model.parameters(), lr=0.01)
new_scaler = GradScaler()
new_model.load_state_dict(checkpoint['model_state_dict'])
new_optimizer.load_state_dict(checkpoint['optimizer_state_dict'])
new_scaler.load_state_dict(checkpoint['scaler_state_dict'])
print("Model restored successfully")
Результат выполнения кода:
"Model restored successfully"
Пример
Обработаем ситуацию с неполным словарём состояния:
import torch
from torch.cuda.amp import GradScaler
torch.manual_seed(0)
scaler = GradScaler()
state_dict = scaler.state_dict()
del state_dict['_scale']
new_scaler = GradScaler()
result = new_scaler.load_state_dict(state_dict, strict=False)
print(result)
Результат выполнения кода:
LoadStateDictReturn(missing_keys=['_scale'], unexpected_keys=[])
Как видим, при использовании параметра strict=False метод загружает доступные ключи и сообщает об отсутствующих.
Смотрите также
-
класс
GradScaler,
который обеспечивает масштабирование градиентов для повышения стабильности обучения -
метод
state_dict,
который возвращает словарь состояния градиентного скалера -
метод
scale,
который масштабирует градиенты перед обратным распространением -
метод
step,
который выполняет шаг оптимизатора с учётом масштабирования