Метод zero_grad
Метод zero_grad класса Optimizer используется для обнуления
градиентов всех параметров, которые были переданы оптимизатору.
Это необходимо делать перед каждым новым шагом оптимизации,
чтобы градиенты не накапливались от предыдущих итераций.
Метод не принимает никаких параметров и не возвращает значений.
Синтаксис
optimizer.zero_grad()
Пример
Рассмотрим базовый пример использования метода zero_grad
в цикле обучения:
import torch
import torch.nn as nn
import torch.optim as optim
model = nn.Linear(5, 1)
optimizer = optim.SGD(model.parameters(), lr=0.01)
for i in range(3):
data = torch.randn(10, 5)
target = torch.randn(10, 1)
optimizer.zero_grad()
output = model(data)
loss = nn.MSELoss()(output, target)
loss.backward()
optimizer.step()
print(f"Step {i+1}, loss: {loss.item():.4f}")
Результат выполнения кода:
Step 1, loss: 1.5084
Step 2, loss: 1.0873
Step 3, loss: 0.9870
Пример
Покажем, что происходит без вызова zero_grad.
Градиенты будут накапливаться, что приведёт к некорректным обновлениям:
import torch
import torch.nn as nn
import torch.optim as optim
model = nn.Linear(2, 1)
optimizer = optim.SGD(model.parameters(), lr=0.1)
data = torch.tensor([[1.0, 2.0]])
target = torch.tensor([[3.0]])
for i in range(2):
# Без вызова zero_grad
output = model(data)
loss = nn.MSELoss()(output, target)
loss.backward()
optimizer.step()
print(f"Step {i+1}, grad norm: {model.weight.grad.norm().item():.4f}")
Результат выполнения кода:
Step 1, grad norm: 7.9702
Step 2, grad norm: 19.4547
Пример
Метод zero_grad можно вызвать с аргументом
set_to_none, который позволяет обнулять градиенты
более эффективно:
import torch
import torch.nn as nn
import torch.optim as optim
model = nn.Sequential(
nn.Linear(3, 5),
nn.ReLU(),
nn.Linear(5, 2),
)
optimizer = optim.Adam(model.parameters(), lr=0.001)
# Обнуление градиентов с установкой в None
optimizer.zero_grad(set_to_none=True)
data = torch.randn(4, 3)
target = torch.randn(4, 2)
output = model(data)
loss = nn.MSELoss()(output, target)
loss.backward()
optimizer.step()
print(f"Loss after step: {loss.item():.4f}")
Результат выполнения кода:
Loss after step: 1.7559
Смотрите также
-
класс
Optimizer,
который является базовым для всех оптимизаторов -
метод
step,
который обновляет параметры модели -
метод
add_param_group,
который добавляет новую группу параметров -
атрибут
param_groups,
который содержит группы параметров оптимизатора