Spec-Zone.ru › PyTorch 1

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. Перед возвратом их пользователю выполняется проверка, чтобы убедиться, что они не использовались в операциях замены на месте, которые изменили их содержимое.

Аргументы также могут быть 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)

© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/1.13/generated/torch.autograd.function.FunctionCtx.save_for_backward.html

Spec-Zone.ru

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