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

Класс LRScheduler

Класс LRScheduler является абстрактным базовым классом для всех планировщиков скорости обучения в PyTorch. Он предоставляет интерфейс для изменения скорости обучения оптимизатора в процессе обучения нейронной сети. Основная идея заключается в том, чтобы адаптивно изменять скорость обучения на протяжении эпох или итераций для улучшения сходимости модели. При создании собственного планировщика необходимо переопределить метод get_lr, который вычисляет новую скорость обучения на каждом шаге.

Синтаксис

class LRScheduler: def __init__(self, optimizer, last_epoch=-1, verbose=False): """ optimizer - оптимизатор, скорость обучения которого нужно изменять last_epoch - индекс последней эпохи (по умолчанию -1) verbose - если True, выводить информацию о изменении скорости обучения """ def step(self, epoch=None): """Обновляет скорость обучения""" def get_last_lr(self): """Возвращает последние вычисленные скорости обучения""" def get_lr(self): """Вычисляет новые скорости обучения (должен быть переопределен)"""

Пример базового использования

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

import torch from torch.optim import SGD from torch.optim.lr_scheduler import LRScheduler class CustomScheduler(LRScheduler): def __init__(self, optimizer, decay_factor=0.1, last_epoch=-1): self.decay_factor = decay_factor super().__init__(optimizer, last_epoch) def get_lr(self): return [base_lr * (1 - self.decay_factor * self.last_epoch) for base_lr in self.base_lrs] model = torch.nn.Linear(10, 1) optimizer = SGD(model.parameters(), lr=0.1) scheduler = CustomScheduler(optimizer, decay_factor=0.01) for epoch in range(5): optimizer.zero_grad() # Симуляция обучения scheduler.step() print(f"Epoch {epoch}: lr = {optimizer.param_groups[0]['lr']:.4f}")

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

Epoch 0: lr = 0.0990 Epoch 1: lr = 0.0980 Epoch 2: lr = 0.0970 Epoch 3: lr = 0.0960 Epoch 4: lr = 0.0950

Пример работы с несколькими группами параметров

Планировщик может управлять скоростью обучения для разных групп параметров независимо:

import torch from torch.optim import SGD from torch.optim.lr_scheduler import LRScheduler class MultiGroupScheduler(LRScheduler): def __init__(self, optimizer, decay_factors, last_epoch=-1): self.decay_factors = decay_factors super().__init__(optimizer, last_epoch) def get_lr(self): return [ base_lr * (1 - factor * self.last_epoch) for base_lr, factor in zip(self.base_lrs, self.decay_factors) ] model = torch.nn.Sequential( torch.nn.Linear(10, 5), torch.nn.Linear(5, 1) ) optimizer = SGD([ {'params': model[0].parameters(), 'lr': 0.1}, {'params': model[1].parameters(), 'lr': 0.01} ]) scheduler = MultiGroupScheduler(optimizer, decay_factors=[0.02, 0.01]) for epoch in range(3): scheduler.step() for i, group in enumerate(optimizer.param_groups): print(f"Epoch {epoch}, group {i}: lr = {group['lr']:.4f}")

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

Epoch 0, group 0: lr = 0.0980 Epoch 0, group 1: lr = 0.0099 Epoch 1, group 0: lr = 0.0960 Epoch 1, group 1: lr = 0.0098 Epoch 2, group 0: lr = 0.0940 Epoch 2, group 1: lr = 0.0097

Пример с методом get_last_lr

Метод get_last_lr позволяет получить последние вычисленные значения скорости обучения:

import torch from torch.optim import SGD from torch.optim.lr_scheduler import LRScheduler class SimpleScheduler(LRScheduler): def __init__(self, optimizer, factor=0.5, last_epoch=-1): self.factor = factor super().__init__(optimizer, last_epoch) def get_lr(self): return [base_lr * (self.factor ** self.last_epoch) for base_lr in self.base_lrs] model = torch.nn.Linear(10, 1) optimizer = SGD(model.parameters(), lr=1.0) scheduler = SimpleScheduler(optimizer, factor=0.5) for epoch in range(3): scheduler.step() print(f"Current lr: {optimizer.param_groups[0]['lr']:.3f}") print(f"Last lr: {scheduler.get_last_lr()[0]:.3f}")

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

Current lr: 0.500 Last lr: 0.500 Current lr: 0.250 Last lr: 0.250 Current lr: 0.125 Last lr: 0.125

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

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