РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
603 of 769 menu

Класс 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,
    который применяет пользовательскую функцию к входным данным
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить