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

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