Метод 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,
который обновляет масштабировщик после шага оптимизации