Метод vjp
Метод vjp класса Function используется для вычисления
векторно-якобиановского произведения (Vector-Jacobian Product).
Этот метод вызывается во время обратного распространения ошибки
и позволяет задать правило вычисления градиентов для пользовательской
функции. Метод принимает выходные данные функции и тензор градиентов
из вышестоящих слоёв, а возвращает кортеж градиентов по входным данным.
Синтаксис
class MyFunction(torch.autograd.Function):
@staticmethod
def forward(ctx, input):
# прямое распространение
return output
@staticmethod
def vjp(ctx, grad_output):
# вычисление градиентов
return grad_input
Метод vjp принимает следующие параметры:
-
ctx- контекстный объект, который сохраняет данные из методаforward -
grad_output- тензор градиентов, поступивший от вышестоящих слоёв
Пример
Давайте создадим пользовательскую функцию, которая вычисляет квадрат
входного тензора, и реализуем метод vjp для вычисления
градиента:
import torch
class SquareFunction(torch.autograd.Function):
@staticmethod
def forward(ctx, input):
ctx.save_for_backward(input)
return input ** 2
@staticmethod
def vjp(ctx, grad_output):
input, = ctx.saved_tensors
grad_input = 2 * input * grad_output
return grad_input
x = torch.tensor([1.0, 2.0, 3.0, 4.0, 5.0], requires_grad=True)
y = SquareFunction.apply(x)
y.sum().backward()
print(x.grad)
Результат выполнения кода:
tensor([2., 4., 6., 8., 10.])
Пример
Рассмотрим более сложную функцию - возведение в куб с сохранением промежуточных вычислений в контексте:
import torch
class CubeFunction(torch.autograd.Function):
@staticmethod
def forward(ctx, input):
ctx.save_for_backward(input)
return input ** 3
@staticmethod
def vjp(ctx, grad_output):
input, = ctx.saved_tensors
grad_input = 3 * input ** 2 * grad_output
return grad_input
torch.manual_seed(0)
x = torch.randn(3, requires_grad=True)
y = CubeFunction.apply(x)
y.sum().backward()
print(x.grad)
Результат выполнения кода:
tensor([0.5327, 0.1885, 0.3232])
Пример
Создадим функцию, которая принимает несколько входных аргументов
и возвращает несколько выходных значений. Метод vjp должен
возвращать градиент для каждого входного аргумента:
import torch
class AddMulFunction(torch.autograd.Function):
@staticmethod
def forward(ctx, a, b):
ctx.save_for_backward(a, b)
return a + b, a * b
@staticmethod
def vjp(ctx, grad_output1, grad_output2):
a, b = ctx.saved_tensors
grad_a = grad_output1 + b * grad_output2
grad_b = grad_output1 + a * grad_output2
return grad_a, grad_b
a = torch.tensor([2.0, 3.0, 4.0], requires_grad=True)
b = torch.tensor([5.0, 6.0, 7.0], requires_grad=True)
sum_res, mul_res = AddMulFunction.apply(a, b)
loss = sum_res.sum() + mul_res.sum()
loss.backward()
print(a.grad)
print(b.grad)
Результат выполнения кода:
tensor([6., 7., 8.])
tensor([3., 4., 5.])