Метод 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])