Метод apply
Статический метод apply класса Function в PyTorch является точкой входа для выполнения пользовательской операции, определенной через наследование от torch.autograd.Function. Он принимает входные тензоры и дополнительные аргументы, вызывает методы forward и setup_context, и возвращает результат, участвуя в графе вычислений для автоматического дифференцирования.
Синтаксис
torch.autograd.Function.apply(*args, **kwargs)
Метод apply принимает произвольное количество позиционных и именованных аргументов, которые передаются в метод forward класса. Возвращает тензор или кортеж тензоров, полученный в результате выполнения прямого прохода.
Пример
Рассмотрим пример создания пользовательской операции, которая вычисляет квадрат входного тензора:
import torch
class Square(torch.autograd.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
x = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
res = Square.apply(x)
print(res)
Результат выполнения кода:
tensor([1., 4., 9.], grad_fn=<SquareBackward>)
Пример
Используем метод apply с несколькими входными тензорами. Создадим операцию, вычисляющую сумму произведений двух тензоров:
import torch
class SumProd(torch.autograd.Function):
@staticmethod
def forward(ctx, a, b):
ctx.save_for_backward(a, b)
return torch.sum(a * b)
@staticmethod
def backward(ctx, grad_output):
a, b = ctx.saved_tensors
return grad_output * b, grad_output * a
t1 = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
t2 = torch.tensor([4.0, 5.0, 6.0], requires_grad=True)
res = SumProd.apply(t1, t2)
print(res)
Результат выполнения кода:
tensor(32., grad_fn=<SumProdBackward>)
Пример
Метод apply может принимать дополнительные аргументы, которые не являются тензорами, но влияют на вычисления:
import torch
class ScaleAdd(torch.autograd.Function):
@staticmethod
def forward(ctx, x, scale, add):
ctx.save_for_backward(x)
ctx.scale = scale
ctx.add = add
return x * scale + add
@staticmethod
def backward(ctx, grad_output):
x, = ctx.saved_tensors
return grad_output * ctx.scale, None, None
x = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
res = ScaleAdd.apply(x, 2.0, 5.0)
print(res)
Результат выполнения кода:
tensor([7., 9., 11.], grad_fn=<ScaleAddBackward>)
Смотрите также
-
класс
Function,
базовый класс для создания пользовательских autograd-операций -
метод
forward,
определяющий прямое распространение в пользовательской функции -
метод
backward,
определяющий обратное распространение градиента -
метод
setup_context,
используемый для сохранения данных между forward и backward