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

Метод 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,
    который выполняет шаг оптимизатора с учётом масштабирования
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить