Метод forward
Метод forward класса Function определяет вычисления,
которые выполняются при прямом проходе в пользовательской операции
автограда. Первым параметром метод принимает контекст ctx,
затем произвольное количество входных тензоров. Метод должен
возвращать один или несколько тензоров - результаты вычислений.
Внутри метода необходимо сохранять данные для обратного прохода
с помощью методов ctx.save_for_backward или
ctx.set_materialize_grads.
Синтаксис
class MyFunction(torch.autograd.Function):
@staticmethod
def forward(ctx, input, param1, param2):
# сохранение данных для backward
ctx.save_for_backward(input)
# вычисления
result = input * param1 + param2
return result
Пример
Создадим операцию, которая умножает тензор на коэффициент и добавляет смещение:
import torch
class ScaleAndShift(torch.autograd.Function):
@staticmethod
def forward(ctx, input, scale, shift):
ctx.save_for_backward(input, scale)
result = input * scale + shift
return result
t = torch.tensor([1.0, 2.0, 3.0, 4.0, 5.0])
res = ScaleAndShift.apply(t, 2.0, 1.0)
print(res)
Результат выполнения кода:
tensor([3., 5., 7., 9., 11.])
Пример
Реализуем операцию возведения в квадрат с сохранением входных данных для обратного прохода:
import torch
class Square(torch.autograd.Function):
@staticmethod
def forward(ctx, input):
ctx.save_for_backward(input)
result = input ** 2
return result
t = torch.tensor([1.0, 2.0, 3.0, 4.0, 5.0], requires_grad=True)
res = Square.apply(t)
print(res)
Результат выполнения кода:
tensor([1., 4., 9., 16., 25.], grad_fn=<SquareBackward>)
Пример
Создадим операцию, которая принимает два тензора и возвращает их поэлементное произведение:
import torch
class Multiply(torch.autograd.Function):
@staticmethod
def forward(ctx, a, b):
ctx.save_for_backward(a, b)
result = a * b
return result
t1 = torch.tensor([1.0, 2.0, 3.0, 4.0, 5.0])
t2 = torch.tensor([5.0, 4.0, 3.0, 2.0, 1.0])
res = Multiply.apply(t1, t2)
print(res)
Результат выполнения кода:
tensor([5., 8., 9., 8., 5.])
Пример
Реализуем операцию, которая обрабатывает двумерный тензор, вычисляя сумму строк:
import torch
class SumRows(torch.autograd.Function):
@staticmethod
def forward(ctx, input):
ctx.save_for_backward(input)
result = input.sum(dim=1)
return result
t = torch.tensor([
[1.0, 2.0, 3.0],
[4.0, 5.0, 6.0],
])
res = SumRows.apply(t)
print(res)
Результат выполнения кода:
tensor([6., 15.])
Пример
Создадим операцию, которая возвращает несколько тензоров:
import torch
class SplitAndScale(torch.autograd.Function):
@staticmethod
def forward(ctx, input, scale):
ctx.save_for_backward(input, scale)
half = input.shape[0] // 2
first = input[:half] * scale
second = input[half:] * scale
return first, second
t = torch.tensor([1.0, 2.0, 3.0, 4.0, 5.0, 6.0])
res1, res2 = SplitAndScale.apply(t, 2.0)
print(res1)
print(res2)
Результат выполнения кода:
tensor([2., 4., 6.])
tensor([8., 10., 12.])
Смотрите также
-
класс
Function,
который является базовым для создания пользовательских операций -
метод
backward,
который определяет логику обратного распространения -
метод
apply,
который вызывает выполнение прямого прохода -
метод
setup_context,
который настраивает контекст для обратного прохода