Функция backward
Функция backward вычисляет градиенты для тензоров,
для которых включено отслеживание операций (requires_grad=True).
Этот метод является основным механизмом обратного распространения
ошибки в PyTorch. Вызов backward на скалярном тензоре
запускает цепочку вычислений градиентов для всех участвовавших
в создании тензора переменных, сохраняя результаты в их атрибуте
grad.
Синтаксис
tensor.backward(gradient=None, retain_graph=False, create_graph=False)
Параметры:
-
gradient- градиент для тензора-приемника, используется для нескалярных тензоров -
retain_graph- сохранять ли граф вычислений после вызова метода -
create_graph- создавать ли граф для вычисления производных высших порядков
Пример
Давайте вычислим градиент функции y = x^2 в точке x = 3:
import torch
x = torch.tensor(3.0, requires_grad=True)
y = x ** 2
y.backward()
print(x.grad)
Результат выполнения кода:
tensor(6.)
Пример
Вычислим градиент для функции с несколькими переменными:
import torch
x = torch.tensor(2.0, requires_grad=True)
y = torch.tensor(3.0, requires_grad=True)
z = x * y + x ** 2
z.backward()
print(x.grad)
print(y.grad)
Результат выполнения кода:
tensor(7.)
tensor(2.)
Пример
Использование backward для нескалярного тензора
с передачей внешнего градиента:
import torch
x = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
y = x ** 2
gradient = torch.tensor([1.0, 0.5, 0.25])
y.backward(gradient)
print(x.grad)
Результат выполнения кода:
tensor([2., 2., 1.5])
Пример
Вычисление градиента с сохранением графа для повторного использования:
import torch
x = torch.tensor(2.0, requires_grad=True)
y = x ** 3
y.backward(retain_graph=True)
print(x.grad)
y.backward()
print(x.grad)
Результат выполнения кода:
tensor(12.)
tensor(24.)
Пример
Вычисление производной второго порядка с помощью
create_graph:
import torch
x = torch.tensor(2.0, requires_grad=True)
y = x ** 3
y.backward(create_graph=True)
grad = x.grad
grad.backward()
print(x.grad)
Результат выполнения кода:
tensor(12.)
Смотрите также
-
функцию
no_grad,
которая отключает вычисление градиентов для оптимизации памяти -
функцию
enable_grad,
которая включает вычисление градиентов в контексте -
функцию
set_grad_enabled,
которая управляет режимом вычисления градиентов глобально -
функцию
grad,
которая вычисляет градиенты для нескольких выходов