Класс Adamax
Класс Adamax реализует алгоритм оптимизации, являющийся расширением метода Adam для случая нормы Lp с p -> ∞. В отличие от Adam, где используется L2-норма градиентов, Adamax применяет бесконечную норму, что делает его более устойчивым при разреженных градиентах. Первым параметром конструктор принимает параметры модели для оптимизации. Вторым параметром можно задать скорость обучения lr. Третьим параметром передаются кортежи betas для коэффициентов затухания моментов. Четвертым параметром задается вес регуляризации weight_decay.
Синтаксис
torch.optim.Adamax(
params,
lr=0.002,
betas=(0.9, 0.999),
eps=1e-8,
weight_decay=0
)
Пример
Создадим простую линейную модель и обучим её с использованием оптимизатора Adamax:
import torch
import torch.nn as nn
torch.manual_seed(0)
model = nn.Linear(10, 1)
optimizer = torch.optim.Adamax(model.parameters(), lr=0.002)
print(optimizer)
Результат выполнения кода:
Adamax (
Parameter Group 0
amsgrad: False
betas: (0.9, 0.999)
eps: 1e-08
lr: 0.002
weight_decay: 0
)
Пример
Выполним один шаг оптимизации, используя случайный входной тензор и целевую переменную:
import torch
import torch.nn as nn
torch.manual_seed(0)
model = nn.Linear(10, 1)
optimizer = torch.optim.Adamax(model.parameters(), lr=0.002)
x = torch.randn(5, 10)
y = torch.randn(5, 1)
criterion = nn.MSELoss()
optimizer.zero_grad()
output = model(x)
loss = criterion(output, y)
loss.backward()
optimizer.step()
print(loss.item())
Результат выполнения кода:
0.9806023836135864
Пример
Продемонстрируем использование параметра weight_decay для L2-регуляризации:
import torch
import torch.nn as nn
torch.manual_seed(0)
model = nn.Linear(10, 1)
optimizer = torch.optim.Adamax(
model.parameters(),
lr=0.002,
weight_decay=1e-4
)
print(optimizer)
Результат выполнения кода:
Adamax (
Parameter Group 0
amsgrad: False
betas: (0.9, 0.999)
eps: 1e-08
lr: 0.002
weight_decay: 0.0001
)
Смотрите также
-
класс
Adam,
который реализует оптимизацию на основе адаптивных моментов с L2-нормой -
класс
AdamW,
который улучшает регуляризацию за счёт корректной реализации L2-штрафа -
класс
SGD,
который реализует стохастический градиентный спуск с инерцией -
класс
RMSprop,
который использует среднеквадратичное значение градиентов для адаптации скорости обучения