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

Метод state_dict

Метод state_dict класса GradScaler возвращает словарь, содержащий текущее состояние масштабирования. Этот словарь можно сохранить на диск и использовать для восстановления состояния масштабирования с помощью метода load_state_dict. Метод не принимает никаких параметров.

Синтаксис

scaler.state_dict()

Пример с сохранением состояния

Давайте создадим масштабировщик, выполним несколько шагов обучения и сохраним его состояние:

import torch from torch.cuda.amp import GradScaler torch.manual_seed(0) model = torch.nn.Linear(10, 1) optimizer = torch.optim.SGD(model.parameters(), lr=0.01) scaler = GradScaler() data = torch.randn(5, 10) target = torch.randn(5, 1) with torch.cuda.amp.autocast(): output = model(data) loss = torch.nn.functional.mse_loss(output, target) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() state = scaler.state_dict() print(state.keys())

Результат выполнения кода:

dict_keys(['scale', 'growth_factor', 'backoff_factor', 'growth_interval', '_growth_tracker'])

Пример с сохранением и загрузкой

Давайте сохраним состояние масштабировщика в файл и восстановим его в новом масштабировщике:

import torch from torch.cuda.amp import GradScaler torch.manual_seed(0) scaler1 = GradScaler() for i in range(3): model = torch.nn.Linear(10, 1) optimizer = torch.optim.SGD(model.parameters(), lr=0.01) data = torch.randn(5, 10) target = torch.randn(5, 1) with torch.cuda.amp.autocast(): output = model(data) loss = torch.nn.functional.mse_loss(output, target) scaler1.scale(loss).backward() scaler1.step(optimizer) scaler1.update() torch.save(scaler1.state_dict(), 'scaler_state.pt') scaler2 = GradScaler() scaler2.load_state_dict(torch.load('scaler_state.pt')) print(scaler2.state_dict()['scale'])

Результат выполнения кода:

tensor(65536.0)

Пример с продолжением обучения

Давайте загрузим состояние масштабировщика и продолжим обучение с того же места:

import torch from torch.cuda.amp import GradScaler torch.manual_seed(0) model = torch.nn.Linear(10, 1) optimizer = torch.optim.SGD(model.parameters(), lr=0.01) scaler = GradScaler() for i in range(2): data = torch.randn(5, 10) target = torch.randn(5, 1) with torch.cuda.amp.autocast(): output = model(data) loss = torch.nn.functional.mse_loss(output, target) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() torch.save(scaler.state_dict(), 'scaler_state.pt') scaler = GradScaler() scaler.load_state_dict(torch.load('scaler_state.pt')) for i in range(2): data = torch.randn(5, 10) target = torch.randn(5, 1) with torch.cuda.amp.autocast(): output = model(data) loss = torch.nn.functional.mse_loss(output, target) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() print(scaler.state_dict()['scale'])

Результат выполнения кода:

tensor(65536.0)

Смотрите также

  • класс GradScaler,
    который обеспечивает автоматическое масштабирование градиентов
  • метод scale,
    который масштабирует потери перед обратным распространением
  • метод load_state_dict,
    который восстанавливает состояние масштабировщика из словаря
  • метод update,
    который обновляет масштабировщик после шага оптимизации
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить