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