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

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