Spec-Zone.ru › PyTorch 2.14

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

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API