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

Класс Optimizer

Класс Optimizer является базовым для всех оптимизаторов в PyTorch. Он предоставляет интерфейс для обновления параметров модели, используя вычисленные градиенты. При создании объекта оптимизатора ему передаются параметры модели, которые будут обновляться. Класс поддерживает такие важные методы, как step для выполнения шага оптимизации, zero_grad для обнуления градиентов и add_param_group для добавления новых групп параметров.

Основные атрибуты класса: param_groups - список групп параметров с их гиперпараметрами, state - словарь для хранения состояния оптимизатора (например, моментов в SGD с моментом) и defaults - словарь с параметрами оптимизации по умолчанию.

Синтаксис

optimizer = torch.optim.Optimizer(params, defaults) # Обычно используется конкретный подкласс, например: optimizer = torch.optim.SGD(model.parameters(), lr=0.01)

Основные методы

Рассмотрим ключевые методы класса Optimizer на примерах.

Пример

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

import torch import torch.nn as nn # Создаем простую модель model = nn.Linear(10, 5) optimizer = torch.optim.SGD(model.parameters(), lr=0.01) # Обнуляем градиенты optimizer.zero_grad()

Пример

Метод step выполняет один шаг оптимизации, обновляя параметры модели на основе текущих градиентов. Обычно он вызывается после backward:

import torch import torch.nn as nn # Простая модель и оптимизатор model = nn.Linear(10, 5) optimizer = torch.optim.SGD(model.parameters(), lr=0.01) # Создаем фиктивные данные x = torch.randn(1, 10) y = torch.randn(1, 5) # Прямой проход, вычисление потерь, обратный проход output = model(x) loss = nn.MSELoss()(output, y) loss.backward() # Шаг оптимизации optimizer.step()

Пример

Метод add_param_group позволяет динамически добавлять новую группу параметров с собственными гиперпараметрами. Это полезно, например, при fine-tuning, когда для разных слоев нужны разные скорости обучения:

import torch import torch.nn as nn # Создаем модель с двумя слоями model = nn.Sequential( nn.Linear(10, 20), nn.ReLU(), nn.Linear(20, 5) ) # Оптимизатор для первого слоя optimizer = torch.optim.SGD(model[0].parameters(), lr=0.01) # Добавляем группу для второго слоя с другой скоростью обучения optimizer.add_param_group({'params': model[2].parameters(), 'lr': 0.001}) # Теперь оптимизатор управляет двумя группами параметров for group in optimizer.param_groups: print(f"Learning rate: {group['lr']}")

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

Learning rate: 0.01 Learning rate: 0.001

Пример

Методы state_dict и load_state_dict используются для сохранения и загрузки состояния оптимизатора. Это важно для возобновления обучения:

import torch import torch.nn as nn # Создаем модель и оптимизатор model = nn.Linear(10, 5) optimizer = torch.optim.SGD(model.parameters(), lr=0.01) # Выполняем несколько шагов for _ in range(5): x = torch.randn(1, 10) y = torch.randn(1, 5) optimizer.zero_grad() loss = nn.MSELoss()(model(x), y) loss.backward() optimizer.step() # Сохраняем состояние оптимизатора state = optimizer.state_dict() print("State saved") # Создаем новый оптимизатор и загружаем состояние new_optimizer = torch.optim.SGD(model.parameters(), lr=0.01) new_optimizer.load_state_dict(state) print("State loaded")

Атрибуты класса

Атрибут param_groups - это список словарей, каждый из которых содержит параметры и их гиперпараметры (например, скорость обучения, коэффициент момента). Атрибут state хранит внутреннее состояние оптимизатора, например, накопленные моменты. Атрибут defaults содержит значения гиперпараметров по умолчанию:

import torch import torch.nn as nn model = nn.Linear(10, 5) optimizer = torch.optim.SGD(model.parameters(), lr=0.01, momentum=0.9) # Просмотр групп параметров print(optimizer.param_groups) # Просмотр состояния (изначально пусто) print(optimizer.state) # Просмотр параметров по умолчанию print(optimizer.defaults)

Результат выполнения кода (примерный вывод):

[{'params': [Parameter containing: ...], 'lr': 0.01, 'momentum': 0.9, 'dampening': 0, 'weight_decay': 0, 'nesterov': False}] {} {'lr': 0.01, 'momentum': 0.9, 'dampening': 0, 'weight_decay': 0, 'nesterov': False}

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

  • метод step,
    который выполняет обновление параметров модели
  • метод zero_grad,
    который обнуляет градиенты всех параметров
  • метод add_param_group,
    который добавляет новую группу параметров
  • метод state_dict,
    который возвращает состояние оптимизатора в виде словаря
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить