Spec-Zone.ru › PyTorch 2.14

torch.autograd.function.FunctionCtx.mark_non_differentiable

FunctionCtx.mark_non_differentiable(*args) [исходный код]

Помечает выходные данные как недифференцируемые.

Этот метод следует вызывать не более одного раза — в методе setup_context() или forward(), причём все аргументы должны быть выходными данными-тензорами.

Это пометит выходные данные как не требующие вычисления градиентов и повысит эффективность обратного прохода. В backward() всё равно нужно принимать градиент для каждого выходного значения, но он всегда будет нулевым тензором той же формы, что и соответствующее выходное значение.

Например, это используется для индексов, возвращаемых при сортировке. См. пример::
>>> class Func(Function):
>>>     @staticmethod
>>>     def forward(ctx, x):
>>>         sorted, idx = x.sort()
>>>         ctx.mark_non_differentiable(idx)
>>>         ctx.save_for_backward(x, idx)
>>>         return sorted, idx
>>>
>>>     @staticmethod
>>>     @once_differentiable
>>>     def backward(ctx, g1, g2):  # still need to accept g2
>>>         x, idx = ctx.saved_tensors
>>>         grad_input = torch.zeros_like(x)
>>>         grad_input.index_add_(0, idx, g1)
>>>         return grad_input

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

Spec-Zone.ru

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