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

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