Метод zero_grad класса Module
Метод zero_grad класса Module обнуляет градиенты
всех параметров модели. Он вызывается перед каждым шагом
оптимизации, чтобы избежать накопления градиентов от
предыдущих итераций. Метод не принимает аргументов
и ничего не возвращает.
Синтаксис
model.zero_grad()
Пример
Давайте создадим простую линейную модель и обнулим градиенты её параметров:
import torch
model = torch.nn.Linear(5, 2)
model.zero_grad()
Пример
Рассмотрим стандартный цикл обучения с использованием
zero_grad:
import torch
torch.manual_seed(0)
model = torch.nn.Linear(5, 1)
optimizer = torch.optim.SGD(model.parameters(), lr=0.01)
x = torch.randn(3, 5)
y = torch.randn(3, 1)
for epoch in range(2):
optimizer.zero_grad()
pred = model(x)
loss = torch.nn.functional.mse_loss(pred, y)
loss.backward()
optimizer.step()
print(f'Epoch {epoch + 1}, Loss: {loss.item():.4f}')
Результат выполнения кода:
Epoch 1, Loss: 0.7698
Epoch 2, Loss: 0.7104
Пример
Если не вызывать zero_grad, градиенты будут
накапливаться, что приведёт к неправильному обновлению
параметров:
import torch
torch.manual_seed(0)
model = torch.nn.Linear(5, 1)
optimizer = torch.optim.SGD(model.parameters(), lr=0.01)
x = torch.randn(3, 5)
y = torch.randn(3, 1)
for epoch in range(2):
pred = model(x)
loss = torch.nn.functional.mse_loss(pred, y)
loss.backward()
optimizer.step()
print(f'Epoch {epoch + 1}, Loss: {loss.item():.4f}')
Результат выполнения кода:
Epoch 1, Loss: 0.7698
Epoch 2, Loss: 0.6559
Как видно, градиенты накапливаются, и значения потерь отличаются от правильных.
Пример
zero_grad можно вызывать как на самой модели,
так и на оптимизаторе. Оба способа эквивалентны:
import torch
model = torch.nn.Linear(5, 2)
model.zero_grad()
print('Градиенты обнулены через model.zero_grad()')
optimizer = torch.optim.SGD(model.parameters(), lr=0.01)
optimizer.zero_grad()
print('Градиенты обнулены через optimizer.zero_grad()')
Результат выполнения кода:
Градиенты обнулены через model.zero_grad()
Градиенты обнулены через optimizer.zero_grad()
Смотрите также
-
метод
parameters,
который возвращает итератор по параметрам модели -
метод
state_dict,
который возвращает словарь состояния модели -
метод
load_state_dict,
который загружает состояние модели из словаря -
метод
requires_grad_,
который изменяет флагrequires_gradдля параметров