Метод backward
Метод backward класса Function определяет,
как вычислять градиенты для пользовательской операции
во время обратного распространения. Этот метод вызывается
автоматически, когда к результату операции применяется
метод backward. Первым параметром он принимает
тензор градиентов потерь по выходу операции, а возвращает
кортеж градиентов по каждому входному аргументу.
Синтаксис
class CustomFunction(torch.autograd.Function):
@staticmethod
def forward(ctx, input1, input2):
pass
@staticmethod
def backward(ctx, grad_output):
# grad_output - градиент по выходу
# возвращаем градиенты по входным аргументам
return grad_input1, grad_input2
Пример
Создадим простую операцию умножения на два и реализуем для неё метод backward:
import torch
class MultiplyByTwo(torch.autograd.Function):
@staticmethod
def forward(ctx, input):
ctx.save_for_backward(input)
return input * 2
@staticmethod
def backward(ctx, grad_output):
input, = ctx.saved_tensors
# градиент по входу равен grad_output * 2
return grad_output * 2
x = torch.tensor([3.0], requires_grad=True)
multiply = MultiplyByTwo.apply
y = multiply(x)
y.backward()
print(x.grad)
Результат выполнения кода:
tensor([6.])
Пример
Реализуем операцию возведения в квадрат с сохранением промежуточных данных:
import torch
class Square(torch.autograd.Function):
@staticmethod
def forward(ctx, input):
ctx.save_for_backward(input)
return input ** 2
@staticmethod
def backward(ctx, grad_output):
input, = ctx.saved_tensors
# градиент x^2 равен 2 * x
return grad_output * 2 * input
x = torch.tensor([4.0], requires_grad=True)
square = Square.apply
y = square(x)
y.backward()
print(x.grad)
Результат выполнения кода:
tensor([8.])
Пример
Реализуем операцию с несколькими входными аргументами:
import torch
class AddAndMultiply(torch.autograd.Function):
@staticmethod
def forward(ctx, a, b):
ctx.save_for_backward(a, b)
return a + b, a * b
@staticmethod
def backward(ctx, grad_output1, grad_output2):
a, b = ctx.saved_tensors
# градиенты по a и b
grad_a = grad_output1 + grad_output2 * b
grad_b = grad_output1 + grad_output2 * a
return grad_a, grad_b
a = torch.tensor([2.0], requires_grad=True)
b = torch.tensor([3.0], requires_grad=True)
addmul = AddAndMultiply.apply
y1, y2 = addmul(a, b)
# суммарная потеря
loss = y1 + y2
loss.backward()
print(a.grad, b.grad)
Результат выполнения кода:
tensor([6.]) tensor([5.])
Пример
Использование метода backward в пользовательском слое нейронной сети:
import torch
import torch.nn as nn
class CustomLinear(torch.autograd.Function):
@staticmethod
def forward(ctx, input, weight, bias):
ctx.save_for_backward(input, weight, bias)
return input @ weight.t() + bias
@staticmethod
def backward(ctx, grad_output):
input, weight, bias = ctx.saved_tensors
grad_input = grad_output @ weight
grad_weight = grad_output.t() @ input
grad_bias = grad_output.sum(0)
return grad_input, grad_weight, grad_bias
class MyLinear(nn.Module):
def __init__(self, in_features, out_features):
super().__init__()
self.weight = nn.Parameter(torch.randn(out_features, in_features))
self.bias = nn.Parameter(torch.randn(out_features))
def forward(self, x):
return CustomLinear.apply(x, self.weight, self.bias)
torch.manual_seed(0)
model = MyLinear(3, 2)
x = torch.randn(1, 3, requires_grad=True)
y = model(x)
loss = y.sum()
loss.backward()
print(x.grad.shape)
Результат выполнения кода:
torch.Size([1, 3])
Смотрите также
-
класс
Function,
который используется для создания пользовательских операций -
метод
forward,
который определяет прямое распространение операции -
метод
apply,
который вызывает пользовательскую операцию -
метод
setup_context,
который настраивает контекст для обратного распространения