Атрибут retains_grad
Атрибут retains_grad класса Tensor
указывает, должен ли тензор сохранять свои градиенты
после того, как был вызван метод backward.
По умолчанию этот атрибут установлен в False
для всех тензоров.
Когда вы вызываете backward на конечном тензоре,
градиенты вычисляются для всех листьев графа вычислений,
для которых requires_grad установлен в True.
После этого градиенты нелистовых тензоров обычно удаляются
для экономии памяти. Атрибут retains_grad позволяет
изменить это поведение.
Синтаксис
tensor.retains_grad
Атрибут можно проверить, но нельзя установить напрямую.
Для его включения используется метод retain_grad.
Пример
Давайте создадим простой тензор с градиентом и проверим значение атрибута retains_grad:
import torch
t = torch.tensor([1., 2., 3.], requires_grad=True)
print(t.retains_grad)
Результат выполнения кода:
False
Пример
Давайте создадим два тензора с градиентами и выполним операцию умножения:
import torch
t1 = torch.tensor([1., 2., 3.], requires_grad=True)
t2 = torch.tensor([4., 5., 6.], requires_grad=True)
res = t1 * t2
print(res.requires_grad)
print(res.retains_grad)
Результат выполнения кода:
True
False
Хотя результат операции имеет requires_grad=True,
атрибут retains_grad остаётся False.
Пример
Теперь применим метод retain_grad к промежуточному тензору
и посмотрим на его градиент после backward:
import torch
t1 = torch.tensor([1., 2., 3.], requires_grad=True)
t2 = torch.tensor([4., 5., 6.], requires_grad=True)
res = t1 * t2
res.retain_grad()
loss = res.sum()
loss.backward()
print(res.retains_grad)
print(res.grad)
Результат выполнения кода:
True
tensor([1., 1., 1.])
После вызова retain_grad атрибут становится True,
и градиент промежуточного тензора сохраняется.
Пример
Давайте сравним поведение с сохранённым и без сохранения градиента для промежуточного тензора:
import torch
t1 = torch.tensor([1., 2., 3.], requires_grad=True)
t2 = torch.tensor([4., 5., 6.], requires_grad=True)
res_no_retain = t1 * t2
loss = res_no_retain.sum()
loss.backward()
print("Without retain_grad:")
print(res_no_retain.retains_grad)
print(res_no_retain.grad)
Результат выполнения кода:
Without retain_grad:
False
None
Без вызова retain_grad градиент промежуточного тензора
не сохраняется и имеет значение None.
Смотрите также
-
метод
retain_grad,
который включает сохранение градиента для тензора -
атрибут
grad,
который содержит вычисленные градиенты тензора -
атрибут
requires_grad,
который указывает, нужно ли вычислять градиенты для тензора -
метод
backward,
который вычисляет градиенты в графе вычислений