Класс 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,
который сохраняет состояние планировщика