Класс Function
Класс Function из модуля torch.autograd предназначен для создания
пользовательских операций, которые могут участвовать в графе вычислений
и поддерживать обратное распространение ошибки. Он используется в тех
случаях, когда стандартные операции PyTorch не подходят, и требуется
определить собственное прямое и обратное распространение.
Для создания пользовательской операции необходимо унаследоваться от
класса Function и переопределить статические методы
forward и backward. Метод forward выполняет
прямой проход, а backward вычисляет градиенты для входных данных
во время обратного распространения.
Важно отметить, что класс Function не является модулем
(не наследник nn.Module) и предназначен для построения
низкоуровневых операций, которые могут быть использованы в любом месте
графа вычислений.
Синтаксис
class CustomFunction(torch.autograd.Function):
@staticmethod
def forward(ctx, input, param):
# Сохраняем данные для backward
ctx.save_for_backward(input)
ctx.param = param
# Выполняем прямое преобразование
output = input + param
return output
@staticmethod
def backward(ctx, grad_output):
# Восстанавливаем сохранённые данные
input, = ctx.saved_tensors
param = ctx.param
# Вычисляем градиенты
grad_input = grad_output
grad_param = grad_output.sum()
return grad_input, grad_param
Пример
Давайте создадим пользовательскую операцию, которая прибавляет к каждому элементу тензора некоторое число:
import torch
class AddConstant(torch.autograd.Function):
@staticmethod
def forward(ctx, input, constant):
ctx.save_for_backward(input)
ctx.constant = constant
output = input + constant
return output
@staticmethod
def backward(ctx, grad_output):
# Производная по входу равна 1
grad_input = grad_output
# Производная по константе равна сумме градиентов
grad_constant = grad_output.sum()
return grad_input, grad_constant
# Создаём тензор и применяем операцию
t = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
res = AddConstant.apply(t, 5.0)
print(res)
Результат выполнения кода:
tensor([6., 7., 8.], grad_fn=<AddConstantBackward>)
Обратите внимание на grad_fn, который указывает, что тензор
получен с помощью нашей пользовательской функции.
Пример
Теперь вызовем обратное распространение для вычисления градиентов:
import torch
class AddConstant(torch.autograd.Function):
@staticmethod
def forward(ctx, input, constant):
ctx.save_for_backward(input)
ctx.constant = constant
output = input + constant
return output
@staticmethod
def backward(ctx, grad_output):
grad_input = grad_output
grad_constant = grad_output.sum()
return grad_input, grad_constant
t = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
res = AddConstant.apply(t, 5.0)
# Вычисляем градиенты
loss = res.sum()
loss.backward()
print(t.grad)
Результат выполнения кода:
tensor([1., 1., 1.])
Градиент для каждого элемента входного тензора равен единице,
что соответствует производной функции x + const по x.
Пример
Рассмотрим более сложный пример - операцию возведения в квадрат:
import torch
class Square(torch.autograd.Function):
@staticmethod
def forward(ctx, input):
ctx.save_for_backward(input)
output = input ** 2
return output
@staticmethod
def backward(ctx, grad_output):
input, = ctx.saved_tensors
# Производная от x^2 равна 2*x
grad_input = 2 * input * grad_output
return grad_input
t = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
res = Square.apply(t)
print(res)
loss = res.sum()
loss.backward()
print(t.grad)
Результат выполнения кода:
tensor([1., 4., 9.], grad_fn=<SquareBackward>)
tensor([2., 4., 6.])
Градиенты вычислены корректно: для элементов 1, 2, 3
градиенты равны 2, 4, 6 соответственно.
Важные особенности
При создании пользовательской функции необходимо помнить о следующих моментах:
-
Метод
forwardдолжен быть статическим и принимать первым аргументом контекстctx. -
Для сохранения данных между прямым и обратным проходами
используйте методы
ctx.save_for_backward(для тензоров) или просто сохраняйте атрибуты вctx(для скаляров). -
Метод
backwardдолжен возвращать ровно столько градиентов, сколько было входных аргументов уforward(не считаяctx). -
Для применения функции используйте метод
apply, который автоматически создаётся для каждого класса-наследника.
Смотрите также
-
класс
Function,
который является основой для создания пользовательских операций -
метод
forward,
который определяет прямое преобразование в пользовательской функции -
метод
backward,
который вычисляет градиенты для обратного распространения -
метод
apply,
который применяет пользовательскую функцию к входным данным