РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
617 of 769 menu

Функция 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,
    которая вычисляет градиенты для нескольких выходов
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить