Метод step
Метод step класса GradScaler выполняет шаг оптимизации для переданного оптимизатора, учитывая масштабирование градиентов, выполненное ранее с помощью метода scale. В процессе работы метод сначала отменяет масштабирование градиентов (unscale), а затем вызывает step оптимизатора. Если во время обратного распространения обнаружены переполнения (инфы или наны) в масштабированных градиентах, метод пропускает обновление параметров, чтобы избежать повреждения модели. Метод возвращает значение, указывающее, был ли выполнен шаг оптимизации.
Синтаксис
scaler.step(optimizer, *args, **kwargs)
Метод принимает следующие параметры:
-
optimizer- оптимизатор PyTorch, для которого выполняется шаг оптимизации. -
*argsи**kwargs- дополнительные позиционные и именованные аргументы, которые будут переданы в методstepоптимизатора.
Метод возвращает булево значение True, если шаг оптимизации был выполнен, и False, если шаг был пропущен из-за переполнения градиентов.
Пример использования
Давайте рассмотрим базовый пример использования метода step в процессе обучения модели с использованием автоматического смешанного обучения:
import torch
import torch.nn as nn
torch.manual_seed(0)
class SimpleModel(nn.Module):
def __init__(self):
super().__init__()
self.fc = nn.Linear(10, 1)
def forward(self, x):
return self.fc(x)
model = SimpleModel()
optimizer = torch.optim.SGD(model.parameters(), lr=0.01)
scaler = torch.amp.GradScaler()
loss_fn = nn.MSELoss()
data = torch.randn(5, 10)
target = torch.randn(5, 1)
optimizer.zero_grad()
with torch.amp.autocast('cpu'):
output = model(data)
loss = loss_fn(output, target)
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
print("Optimization step completed")
Результат выполнения кода:
"Optimization step completed"
Пример с проверкой выполнения шага
Метод step возвращает булево значение, которое можно использовать для контроля выполнения шага оптимизации:
import torch
import torch.nn as nn
torch.manual_seed(1)
class SimpleModel(nn.Module):
def __init__(self):
super().__init__()
self.fc = nn.Linear(10, 1)
def forward(self, x):
return self.fc(x)
model = SimpleModel()
optimizer = torch.optim.SGD(model.parameters(), lr=0.01)
scaler = torch.amp.GradScaler()
loss_fn = nn.MSELoss()
data = torch.randn(5, 10)
target = torch.randn(5, 1)
optimizer.zero_grad()
with torch.amp.autocast('cpu'):
output = model(data)
loss = loss_fn(output, target)
scaler.scale(loss).backward()
step_performed = scaler.step(optimizer)
scaler.update()
print(f"Step performed: {step_performed}")
Результат выполнения кода:
"Step performed: True"
Пример с пропуском шага
В случае, если во время обратного распространения обнаруживаются градиенты с переполнением, метод step пропускает обновление параметров и возвращает False:
import torch
import torch.nn as nn
torch.manual_seed(2)
class SimpleModel(nn.Module):
def __init__(self):
super().__init__()
self.fc = nn.Linear(10, 1)
self.weight = nn.Parameter(torch.ones(1) * 1e10)
def forward(self, x):
x = self.fc(x)
return x * self.weight
model = SimpleModel()
optimizer = torch.optim.SGD(model.parameters(), lr=0.01)
scaler = torch.amp.GradScaler()
loss_fn = nn.MSELoss()
data = torch.randn(5, 10)
target = torch.randn(5, 1)
optimizer.zero_grad()
with torch.amp.autocast('cpu'):
output = model(data)
loss = loss_fn(output, target)
scaler.scale(loss).backward()
step_performed = scaler.step(optimizer)
scaler.update()
print(f"Step performed: {step_performed}")
Результат выполнения кода:
"Step performed: False"
Смотрите также
-
класс
GradScaler,
который управляет градиентным масштабированием -
метод
scale,
который масштабирует потери перед обратным распространением -
метод
update,
который обновляет коэффициент масштабирования -
метод
unscale_,
который отменяет масштабирование градиентов