torch.autograd.function.FunctionCtx.save_for_backward
-
FunctionCtx.save_for_backward(*tensors)[source] -
Сохраняет заданные тензоры для последующего вызова
backward().save_for_backwardдолжен вызываться не более одного раза, только изнутри методаforward(), и только с тензорами.Все тензоры, предназначенные для использования в обратном проходе, должны быть сохранены с помощью
save_for_backward(а не напрямую вctx) для предотвращения некорректных градиентов и утечек памяти, а также для возможности применения сохраняемых хуков тензоров. См.torch.autograd.graph.saved_tensors_hooks.Обратите внимание, что если промежуточные тензоры, тензоры, которые не являются ни входами, ни выходами
forward(), сохраняются для обратного прохода, ваш пользовательский Функция может не поддерживать двойной обратный проход. Пользовательские функции, которые не поддерживают двойной обратный проход, должны украсить свой методbackward()с@once_differentiable, чтобы выполнение двойного обратного прохода вызывало ошибку. Если вы хотите поддерживать двойной обратный проход, вы можете либо пересчитать промежуточные значения на основе входов во время обратного прохода, либо вернуть промежуточные значения в качестве выходов пользовательской функции. Более подробную информацию см. в руководстве по двойному обратному проходу.В
backward(), сохраненные тензоры можно получить через атрибутsaved_tensors. Перед возвращением их пользователю проверяется, не были ли они использованы в каком-либо операторе in-place, который изменил их содержимое.Аргументы также могут быть
None. Это бесполезная операция.Более подробную информацию о том, как использовать этот метод, см. в Расширение torch.autograd.
- Example::
-
>>> class Func(Function): >>> @staticmethod >>> def forward(ctx, x: torch.Tensor, y: torch.Tensor, z: int): >>> w = x * z >>> out = x * y + y * z + w * y >>> ctx.save_for_backward(x, y, w, out) >>> ctx.z = z # z is not a tensor >>> return out >>> >>> @staticmethod >>> @once_differentiable >>> def backward(ctx, grad_out): >>> x, y, w, out = ctx.saved_tensors >>> z = ctx.z >>> gx = grad_out * (y + y * z) >>> gy = grad_out * (x + z + w) >>> gz = None >>> return gx, gy, gz >>> >>> a = torch.tensor(1., requires_grad=True, dtype=torch.double) >>> b = torch.tensor(2., requires_grad=True, dtype=torch.double) >>> c = 4 >>> d = Func.apply(a, b, c)
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/generated/torch.autograd.function.FunctionCtx.save_for_backward.html