Spec-Zone.ru › PyTorch 1

torch.autograd.function.FunctionCtx.mark_dirty

FunctionCtx.mark_dirty(*args) [source]

Помечает заданные тензоры как изменённые в операциях на месте.

Это должно вызываться не более одного раза, только изнутри forward() метода, и все аргументы должны быть входными.

Каждый тензор, который был изменён на месте в вызове forward() должен быть передан в эту функцию, чтобы обеспечить корректность наших проверок. Не имеет значения, вызывается ли функция до или после изменения.

Примеры::
>>> class Inplace(Function):
>>>     @staticmethod
>>>     def forward(ctx, x):
>>>         x_npy = x.numpy() # x_npy shares storage with x
>>>         x_npy += 1
>>>         ctx.mark_dirty(x)
>>>         return x
>>>
>>>     @staticmethod
>>>     @once_differentiable
>>>     def backward(ctx, grad_output):
>>>         return grad_output
>>>
>>> a = torch.tensor(1., requires_grad=True, dtype=torch.double).clone()
>>> b = a * a
>>> Inplace.apply(a)  # This would lead to wrong gradients!
>>>                   # but the engine would not know unless we mark_dirty
>>> b.backward() # RuntimeError: one of the variables needed for gradient
>>>              # computation has been modified by an inplace operation

© 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.mark_dirty.html

Spec-Zone.ru

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