torch.autograd.function.FunctionCtx.save_for_backward
-
FunctionCtx.save_for_backward(*tensors)[исходный код] -
Сохранить указанные тензоры для последующего вызова
backward().save_for_backwardследует вызывать не более одного раза — в методеsetup_context()илиforward()— и передавать ему только тензоры.Все тензоры, которые будут использоваться при обратном проходе, следует сохранять с помощью
save_for_backward(а не непосредственно вctx), чтобы предотвратить некорректное вычисление градиентов и утечки памяти, а также обеспечить возможность применения хуков для сохранённых тензоров. См.torch.autograd.graph.saved_tensors_hooks. Подробнее см. в разделе Расширение torch.autograd.Обратите внимание: если для обратного прохода сохраняются промежуточные тензоры, то есть тензоры, которые не являются входами или выходами
forward(), ваша пользовательская функция может не поддерживать двойное дифференцирование. Для пользовательских функций, не поддерживающих двойное дифференцирование, следует пометить их методbackward()декоратором@once_differentiable, чтобы при попытке выполнить двойное дифференцирование возникала ошибка. Если вы хотите поддерживать двойное дифференцирование, можно либо повторно вычислять промежуточные тензоры на основе входных данных во время обратного прохода, либо возвращать промежуточные тензоры в качестве выходов пользовательской функции. Подробнее см. в руководстве по двойному дифференцированию.В
backward()доступ к сохранённым тензорам осуществляется через атрибутsaved_tensors. Перед возвратом пользователю проверяется, не использовались ли они в операциях на месте, изменивших их содержимое.Аргументы также могут быть
None. В этом случае ничего не происходит.Подробнее о том, как использовать этот метод, см. в разделе Расширение torch.autograd.
Пример:
>>> 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)
© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/generated/torch.autograd.function.FunctionCtx.save_for_backward.html