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

Метод jvp класса Function

Метод jvp класса Function в PyTorch предназначен для вычисления произведения якобиана (матрицы частных производных) на касательный вектор. Этот метод используется в системе автоматического дифференцирования для реализации режима прямого распространения (forward-mode AD). Метод принимает на вход те же аргументы, что и метод forward, а также касательные векторы для каждого входного тензора. Возвращает метод кортеж из выходных тензоров и касательных векторов для них.

Синтаксис

@staticmethod def jvp(ctx, *inputs, **kwargs): # вычисление прямого распространения # и касательных векторов return outputs, tangent_outputs

Параметры метода:

  • ctx - контекстный объект для сохранения состояния между методами forward и backward
  • *inputs - входные тензоры, переданные в метод apply
  • **kwargs - дополнительные именованные аргументы, переданные в метод apply

Возвращаемое значение - кортеж из двух элементов:

  • результат прямого прохода (такой же, как в forward)
  • касательные векторы для выходных тензоров

Пример

Создадим простую операцию возведения в квадрат и определим для неё метод jvp:

import torch from torch.autograd import Function class Square(Function): @staticmethod def forward(ctx, x): ctx.save_for_backward(x) return x ** 2 @staticmethod def backward(ctx, grad_output): x, = ctx.saved_tensors return grad_output * 2 * x @staticmethod def jvp(ctx, *inputs, **kwargs): x, tangent_x = inputs output = x ** 2 tangent_output = 2 * x * tangent_x return output, tangent_output torch.manual_seed(0) x = torch.tensor([2.0, 3.0, 4.0], requires_grad=True) tangent = torch.tensor([1.0, 0.5, 2.0]) res, tangent_res = Square.apply(x, tangent) print(res) print(tangent_res)

Результат выполнения кода:

tensor([ 4., 9., 16.], grad_fn=<SquareBackward>) tensor([4., 3., 16.])

Пример

Реализуем операцию линейного преобразования с сохранением касательных векторов:

import torch from torch.autograd import Function class LinearTransform(Function): @staticmethod def forward(ctx, x, weight, bias): ctx.save_for_backward(x, weight, bias) return torch.matmul(x, weight) + bias @staticmethod def backward(ctx, grad_output): x, weight, bias = ctx.saved_tensors grad_x = torch.matmul(grad_output, weight.t()) grad_weight = torch.matmul(x.t(), grad_output) grad_bias = grad_output.sum(0) return grad_x, grad_weight, grad_bias @staticmethod def jvp(ctx, *inputs, **kwargs): x, tangent_x = inputs[0], inputs[1] weight, tangent_weight = inputs[2], inputs[3] bias, tangent_bias = inputs[4], inputs[5] output = torch.matmul(x, weight) + bias tangent_output = ( torch.matmul(tangent_x, weight) + torch.matmul(x, tangent_weight) + tangent_bias ) return output, tangent_output torch.manual_seed(0) x = torch.tensor([[1.0, 2.0], [3.0, 4.0]], requires_grad=True) weight = torch.tensor([[0.5, 0.3], [0.7, 0.9]], requires_grad=True) bias = torch.tensor([0.1, 0.2], requires_grad=True) tangent_x = torch.ones_like(x) tangent_weight = torch.zeros_like(weight) tangent_bias = torch.zeros_like(bias) res, tangent_res = LinearTransform.apply( x, tangent_x, weight, tangent_weight, bias, tangent_bias ) print(res) print(tangent_res)

Результат выполнения кода:

tensor([[1.9000, 2.3000], [4.3000, 5.3000]], grad_fn=<LinearTransformBackward>) tensor([[1.2000, 1.2000], [1.2000, 1.2000]])

Пример

Используем метод jvp для вычисления производной по направлению сложной функции:

import torch from torch.autograd import Function class ComplexOp(Function): @staticmethod def forward(ctx, x): ctx.save_for_backward(x) return torch.sin(x) + torch.cos(x) @staticmethod def backward(ctx, grad_output): x, = ctx.saved_tensors return grad_output * (torch.cos(x) - torch.sin(x)) @staticmethod def jvp(ctx, *inputs, **kwargs): x, tangent_x = inputs output = torch.sin(x) + torch.cos(x) tangent_output = (torch.cos(x) - torch.sin(x)) * tangent_x return output, tangent_output torch.manual_seed(0) x = torch.tensor([0.5, 1.0, 1.5], requires_grad=True) tangent = torch.tensor([0.1, 0.2, 0.3]) res, tangent_res = ComplexOp.apply(x, tangent) print(res) print(tangent_res)

Результат выполнения кода:

tensor([1.3570, 1.3818, 1.0688], grad_fn=<ComplexOpBackward>) tensor([0.0836, 0.0906, 0.1199])

Смотрите также

  • класс Function,
    основной класс для создания пользовательских операций
  • метод forward,
    выполняет прямое распространение операции
  • метод backward,
    вычисляет градиенты для обратного распространения
  • метод vjp,
    вычисляет произведение якобиана на вектор для обратного режима
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить