Атрибут grad_fn
Атрибут grad_fn класса Tensor содержит ссылку на объект, который представляет функцию, породившую данный тензор в ходе операций автоматического дифференцирования (autograd). Если тензор создан явно пользователем (например, через torch.tensor), то этот атрибут равен None. Для тензоров, полученных в результате математических операций, grad_fn указывает на функцию, которая была применена, и используется для вычисления градиентов при обратном распространении.
Синтаксис
t.grad_fn
Пример
Проверим значение атрибута для тензора, созданного с помощью torch.tensor:
import torch
t = torch.tensor([1.0, 2.0, 3.0])
print(t.grad_fn)
Результат выполнения кода:
None
Как видите, для тензора, созданного вручную, градиентная функция отсутствует.
Пример
Теперь создадим тензор в результате операции сложения и посмотрим на grad_fn:
import torch
a = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
b = torch.tensor([4.0, 5.0, 6.0], requires_grad=True)
c = a + b
print(c.grad_fn)
Результат выполнения кода:
<AddBackward0 object at 0x7f8a3c1b9a90>
Атрибут указывает на функцию AddBackward0, которая отвечает за вычисление градиентов для операции сложения.
Пример
Для более сложных операций grad_fn будет отражать последовательность применённых функций:
import torch
x = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
y = x ** 2 + 2 * x + 1
print(y.grad_fn)
Результат выполнения кода:
<AddBackward0 object at 0x7f8a3c1b9d60>
В этом случае grad_fn представляет собой последнюю операцию (сложение), а внутри неё хранятся ссылки на предыдущие операции, формируя граф вычислений.
Пример
Проверим атрибут для тензора, созданного с помощью функции torch.ones с включенным отслеживанием градиентов:
import torch
t = torch.ones(3, requires_grad=True)
print(t.grad_fn)
Результат выполнения кода:
Даже если для тензора установлен requires_grad=True, но он создан явным образом, grad_fn остаётся равен None, так как он не является результатом операции.
Смотрите также
-
атрибут
grad,
который содержит вычисленный градиент -
атрибут
requires_grad,
который определяет, нужно ли отслеживать градиенты для тензора -
атрибут
is_leaf,
который показывает, является ли тензор листовым в графе вычислений -
метод
backward,
который запускает процесс вычисления градиентов