Метод 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])