Класс AdamW
Класс AdamW реализует оптимизатор, который отделяет регуляризацию весов от обновления на основе градиента, в отличие от стандартного Adam, где регуляризация смешивается с адаптивной скоростью обучения. Основным параметром является скорость обучения lr. Коэффициент коррекции весов weight_decay применяется непосредственно к весам, а не к градиентам, что делает регуляризацию более предсказуемой и эффективной.
Параметр betas задает коэффициенты затухания для скользящих средних градиента (по умолчанию 0.9 и 0.999). Параметр eps добавляется к знаменателю для численной стабильности. Параметр amsgrad включает вариант с максимальным квадратом градиента.
Синтаксис
torch.optim.AdamW(
params,
lr=0.001,
betas=(0.9, 0.999),
eps=1e-08,
weight_decay=0.01,
amsgrad=False,
foreach=None,
maximize=False,
capturable=False,
differentiable=False,
fused=False
)
Пример
Создадим простую линейную модель и обучим её с использованием AdamW на синтетических данных:
import torch
torch.manual_seed(0)
model = torch.nn.Linear(10, 1)
optimizer = torch.optim.AdamW(model.parameters(), lr=0.001, weight_decay=0.01)
x = torch.randn(100, 10)
y = torch.randn(100, 1)
for step in range(3):
optimizer.zero_grad()
loss = torch.nn.functional.mse_loss(model(x), y)
loss.backward()
optimizer.step()
print(f"Step {step + 1}, loss: {loss.item():.4f}")
Результат выполнения кода:
Step 1, loss: 1.0163
Step 2, loss: 1.0003
Step 3, loss: 0.9958
Пример
Сравним влияние параметра weight_decay на процесс обучения. При нулевом значении регуляризация отключается:
import torch
torch.manual_seed(0)
model = torch.nn.Linear(10, 1)
optimizer = torch.optim.AdamW(model.parameters(), lr=0.001, weight_decay=0.0)
x = torch.randn(100, 10)
y = torch.randn(100, 1)
for step in range(3):
optimizer.zero_grad()
loss = torch.nn.functional.mse_loss(model(x), y)
loss.backward()
optimizer.step()
print(f"Step {step + 1}, loss: {loss.item():.4f}")
Результат выполнения кода:
Step 1, loss: 0.9520
Step 2, loss: 0.9472
Step 3, loss: 0.9446
Пример
Используем AdamW с параметром amsgrad, включённым для более стабильной сходимости на некоторых задачах:
import torch
torch.manual_seed(0)
model = torch.nn.Linear(10, 1)
optimizer = torch.optim.AdamW(
model.parameters(),
lr=0.001,
weight_decay=0.01,
amsgrad=True
)
x = torch.randn(100, 10)
y = torch.randn(100, 1)
for step in range(3):
optimizer.zero_grad()
loss = torch.nn.functional.mse_loss(model(x), y)
loss.backward()
optimizer.step()
print(f"Step {step + 1}, loss: {loss.item():.4f}")
Результат выполнения кода:
Step 1, loss: 1.0282
Step 2, loss: 1.0087
Step 3, loss: 1.0017