Метод requires_grad_
Метод requires_grad_ класса Tensor изменяет значение атрибута requires_grad тензора. Этот атрибут определяет, будет ли тензор отслеживать все операции, выполняемые над ним, для последующего вычисления градиентов с помощью метода backward. Метод возвращает сам тензор, что позволяет использовать его в цепочках вызовов.
Синтаксис
tensor.requires_grad_(requires_grad=True)
Единственным параметром метода является булево значение requires_grad, которое указывает, нужно ли включать отслеживание градиентов. По умолчанию этот параметр равен True.
Пример
Создадим обычный тензор без отслеживания градиентов и включим эту возможность с помощью метода requires_grad_:
import torch
t = torch.tensor([1.0, 2.0, 3.0])
print(t.requires_grad)
t.requires_grad_(True)
print(t.requires_grad)
Результат выполнения кода:
False
True
Как видно, после применения метода атрибут requires_grad изменился на True.
Пример
Метод часто используется для отключения отслеживания градиентов у тензоров, которые не должны участвовать в процессе обучения, например, для заморозки параметров модели. В этом примере мы выключим отслеживание градиентов:
import torch
t = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
print(t.requires_grad)
t.requires_grad_(False)
print(t.requires_grad)
Результат выполнения кода:
True
False
Теперь тензор не отслеживает операции, и вызов backward на нем или на результатах вычислений с его участием будет невозможен.
Пример
Рассмотрим практический пример, где метод requires_grad_ используется для включения отслеживания градиентов у тензора весов перед вычислением градиентов:
import torch
# Создаем тензор весов без градиента
weights = torch.ones(3)
print("До:", weights.requires_grad)
# Включаем отслеживание градиентов
weights.requires_grad_(True)
print("После:", weights.requires_grad)
# Выполняем операцию
x = torch.tensor([1.0, 2.0, 3.0])
loss = (weights * x).sum()
# Вычисляем градиенты
loss.backward()
print(weights.grad)
Результат выполнения кода:
До: False
После: True
tensor([1., 2., 3.])
После включения отслеживания градиентов мы смогли вычислить градиент weights.grad.
Смотрите также
-
атрибут
requires_grad,
который показывает, отслеживает ли тензор операции для градиентов -
атрибут
grad,
который содержит вычисленный градиент тензора -
атрибут
grad_fn,
который указывает функцию, создавшую тензор -
метод
backward,
который вычисляет градиенты для тензора