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

Метод setup_context

Метод setup_context является частью класса Function из модуля torch.autograd. Он используется для сохранения состояния или промежуточных результатов вычислений, которые потребуются при выполнении обратного прохода (метод backward). Этот метод вызывается автоматически в процессе выполнения прямого прохода, определённого в методе forward.

Основное назначение setup_context - отделить логику сохранения данных для градиентов от самой логики прямого вычисления. Это позволяет сделать код более чистым и организованным, особенно при работе с сложными операциями. Метод принимает на вход три параметра: контекст, входные аргументы и выходные данные прямого прохода.

Синтаксис

class MyFunction(torch.autograd.Function): @staticmethod def forward(ctx, *args, **kwargs): # Прямой проход output = ... return output @staticmethod def setup_context(ctx, inputs, output): # Сохранение данных для backward ctx.save_for_backward(*inputs) # Или сохранение других данных ctx.some_data = ... @staticmethod def backward(ctx, grad_output): # Обратный проход с использованием сохранённых данных ...

Пример

Рассмотрим простую операцию умножения на два. В методе setup_context мы сохраняем входной тензор, чтобы в методе backward использовать его для вычисления градиента:

import torch class MulByTwo(torch.autograd.Function): @staticmethod def forward(ctx, x): return x * 2 @staticmethod def setup_context(ctx, inputs, output): x, = inputs ctx.save_for_backward(x) @staticmethod def backward(ctx, grad_output): x, = ctx.saved_tensors # Градиент по x = grad_output * 2 return grad_output * 2 x = torch.tensor([1.0, 2.0, 3.0], requires_grad=True) res = MulByTwo.apply(x) res.sum().backward() print(x.grad)

Результат выполнения кода:

tensor([2., 2., 2.])

Пример

В этом примере мы сохраняем не только входной тензор, но и промежуточный результат. Метод setup_context даёт возможность сохранить любые данные, которые могут понадобиться при обратном проходе:

import torch class PowerFunction(torch.autograd.Function): @staticmethod def forward(ctx, x, exponent): ctx.exponent = exponent # Сохраняем как атрибут return x ** exponent @staticmethod def setup_context(ctx, inputs, output): x, exponent = inputs # Сохраняем входной тензор и результат ctx.save_for_backward(x, output) ctx.exponent = exponent @staticmethod def backward(ctx, grad_output): x, output = ctx.saved_tensors exponent = ctx.exponent # Градиент: exponent * x^(exponent-1) grad_x = grad_output * exponent * (x ** (exponent - 1)) return grad_x, None x = torch.tensor([2.0, 3.0, 4.0], requires_grad=True) res = PowerFunction.apply(x, 3) res.sum().backward() print(x.grad)

Результат выполнения кода:

tensor([12., 27., 48.])

Пример

Метод setup_context особенно полезен, когда нужно сохранить большие тензоры для обратного прохода, но при этом избежать их повторного вычисления:

import torch class ComplexFunction(torch.autograd.Function): @staticmethod def forward(ctx, x, y): # Сложное промежуточное вычисление intermediate = x * y + x output = intermediate.sum() return output @staticmethod def setup_context(ctx, inputs, output): x, y = inputs # Сохраняем все нужные данные ctx.save_for_backward(x, y) ctx.intermediate = x * y + x @staticmethod def backward(ctx, grad_output): x, y = ctx.saved_tensors intermediate = ctx.intermediate # Градиенты по x и y grad_x = grad_output * (y + 1) grad_y = grad_output * x return grad_x, grad_y torch.manual_seed(0) x = torch.randn(3, requires_grad=True) y = torch.randn(3, requires_grad=True) res = ComplexFunction.apply(x, y) res.backward() print(x.grad) print(y.grad)

Результат выполнения кода:

tensor([1.2050, 0.1282, 0.1177]) tensor([-1.2179, -0.5489, -1.2202])

Смотрите также

  • класс Function,
    который является базовым для создания пользовательских операций автограда
  • метод forward,
    который выполняет прямой проход операции
  • метод backward,
    который вычисляет градиенты для обратного распространения
  • метод apply,
    который используется для вызова пользовательской операции
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить