NestedIOFunction
-
class torch.autograd.function.NestedIOFunction(*args, **kwargs)[исходный код] -
Этот класс существует только для обеспечения обратной совместимости. Для новых случаев использования вместо него используйте
Function.-
backward(*gradients)[исходный код] -
Общая вспомогательная функция обратного прохода.
- Тип возвращаемого значения:
-
backward_extended(*grad_output)[исходный код] -
Пользовательская функция обратного прохода.
-
forward(*args)[исходный код] -
Общая вспомогательная функция прямого прохода.
- Тип возвращаемого значения:
-
forward_extended(*input)[исходный код] -
Пользовательская функция прямого прохода.
-
static jvp(ctx, *grad_inputs)[исходный код] -
Задайте формулу дифференцирования операции с помощью автоматического дифференцирования в прямом режиме.
Эту функцию необходимо переопределить во всех подклассах. Она должна принимать контекст
ctxв качестве первого аргумента, а затем столько же входных данных, сколько получила функцияforward()(для нетензорных входных данных функции прямого прохода будет передано None), и возвращать столько же тензоров, сколько было выходных данных уforward(). Каждый аргумент представляет собой градиент по отношению к заданному входному значению, а каждое возвращаемое значение должно быть градиентом по отношению к соответствующему выходному значению. Если выходное значение не является тензором или функция не дифференцируема по этому выходному значению, для соответствующего входного градиента можно передать None.Для передачи любых значений из прямого прохода в эту функцию можно использовать объект
ctx.- Тип возвращаемого значения:
-
mark_dirty(*args, **kwargs)[исходный код] -
См.
Function.mark_dirty().
-
mark_non_differentiable(*args, **kwargs)[исходный код] -
См.
Function.mark_non_differentiable().
-
save_for_backward(*args)[исходный код] -
См.
Function.save_for_backward().
-
save_for_forward(*tensors)[исходный код] -
Сохраните указанные тензоры для последующего вызова
jvp().Вызов
save_for_forwardдопускается не более одного раза — либо в методеsetup_context(), либо в методеforward(); все аргументы должны быть тензорами.В методе
jvp()к сохраненным объектам можно обратиться через атрибутsaved_tensors.В качестве аргументов также можно передать
None. Это не приведет к каким-либо действиям.Подробнее об использовании этого метода см. в разделе Расширение torch.autograd.
Пример:
>>> class Func(torch.autograd.Function): >>> @staticmethod >>> def forward(ctx, x: torch.Tensor, y: torch.Tensor, z: int): >>> ctx.save_for_backward(x, y) >>> ctx.save_for_forward(x, y) >>> ctx.z = z >>> return x * y * z >>> >>> @staticmethod >>> def jvp(ctx, x_t, y_t, _): >>> x, y = ctx.saved_tensors >>> z = ctx.z >>> return z * (y * x_t + x * y_t) >>> >>> @staticmethod >>> def vjp(ctx, grad_out): >>> x, y = ctx.saved_tensors >>> z = ctx.z >>> return z * grad_out * y, z * grad_out * x, None >>> >>> a = torch.tensor(1., requires_grad=True, dtype=torch.double) >>> t = torch.tensor(1., dtype=torch.double) >>> b = torch.tensor(2., requires_grad=True, dtype=torch.double) >>> c = 4 >>> >>> with fwAD.dual_level(): >>> a_dual = fwAD.make_dual(a, t) >>> d = Func.apply(a_dual, b, c)
-
property saved_tensors -
См.
Function.saved_tensors().
-
set_materialize_grads(value)[исходный код] -
Укажите, следует ли материализовать тензоры градиентов. Значение по умолчанию —
True.Этот метод следует вызывать только из метода
setup_context()илиforward().Если задано
True, неопределенные тензоры градиентов перед вызовом методовbackward()иjvp()будут расширены до тензоров, заполненных нулями.Пример:
>>> class SimpleFunc(Function): >>> @staticmethod >>> def forward(ctx, x): >>> return x.clone(), x.clone() >>> >>> @staticmethod >>> @once_differentiable >>> def backward(ctx, g1, g2): >>> return g1 + g2 # No check for None necessary >>> >>> # We modify SimpleFunc to handle non-materialized grad outputs >>> class Func(Function): >>> @staticmethod >>> def forward(ctx, x): >>> ctx.set_materialize_grads(False) >>> ctx.save_for_backward(x) >>> return x.clone(), x.clone() >>> >>> @staticmethod >>> @once_differentiable >>> def backward(ctx, g1, g2): >>> x, = ctx.saved_tensors >>> grad_input = torch.zeros_like(x) >>> if g1 is not None: # We must check for None now >>> grad_input += g1 >>> if g2 is not None: >>> grad_input += g2 >>> return grad_input >>> >>> a = torch.tensor(1., requires_grad=True) >>> b, _ = Func.apply(a) # induces g2 to be undefined
-
set_output_grad_dtype(*dtypes)[исходный код] -
Задайте тип данных градиента для каждого выходного значения этой функции.
Этот метод следует вызывать не более одного раза — либо из метода
setup_context(), либо из методаforward(). Число объявлений должно совпадать с числом возвращаемых значений; каждый аргумент соответствует выходному значению с тем же индексом.Для каждого выходного значения укажите тип данных, в котором обратный проход должен получить его градиент:
- Передайте значение
torch.dtype, и механизм гарантирует, что градиент, передаваемый в обратный проход, будет иметь этот тип данных. Это допустимо только для дифференцируемого тензорного выходного значения. - Передайте
None, и градиент будет передан в обратный проход без изменения его текущего типа данных. Это также единственный допустимый вариант для нетензорного или недифференцируемого выходного значения, у которого нет градиента. - Не вызывайте этот метод (или передайте собственный тип данных выходного значения), и градиент будет передан в обратный проход с типом данных выходного значения; это поведение по умолчанию.
Например:
>>> @staticmethod >>> def forward(ctx, x): >>> t1 = x.sin() >>> t2 = x.cos() >>> t3 = x.tan() >>> ctx.set_output_grad_dtype(torch.float32, t2.dtype, None, None) >>> return t1, t2, t3, "not a tensor"
Это гарантирует, что обратный проход получит градиент
t1в типе данныхfloat32, сохранит поведение по умолчанию для градиентаt2с помощьюt2.dtype, передаст градиентt3без приведения типа с помощьюNoneи используетNoneв качестве заполнителя для последнего нетензорного выходного значения. - Передайте значение
-
static setup_context(ctx, inputs, output)[исходный код] -
Существует два способа определить прямой проход autograd.Function.
Можно:
- Переопределить forward с сигнатурой
forward(ctx, *args, **kwargs). Методsetup_contextне переопределяется. Настройка ctx для обратного прохода выполняется внутриforward. - Переопределить forward с сигнатурой
forward(*args, **kwargs)и переопределитьsetup_context. Настройка ctx для обратного прохода выполняется внутриsetup_context(а не внутриforward).
Подробнее см. в разделе
torch.autograd.Function.forward()и в разделе Расширение torch.autograd.- Тип возвращаемого значения:
- Переопределить forward с сигнатурой
-
static vjp(ctx, *grad_outputs)[исходный код] -
Задайте формулу дифференцирования операции с помощью автоматического дифференцирования в обратном режиме.
Эту функцию необходимо переопределить во всех подклассах. (Определение этой функции эквивалентно определению функции
vjp.)Она должна принимать контекст
ctxв качестве первого аргумента, а затем столько же выходных значений, сколько вернула функцияforward()(для нетензорных выходных значений функции прямого прохода будет передано None), и возвращать столько же тензоров, сколько было входных значений уforward(). Каждый аргумент представляет собой градиент по отношению к заданному выходному значению, а каждое возвращаемое значение должно быть градиентом по отношению к соответствующему входному значению. Если входное значение не является тензором или является тензором, для которого не требуются градиенты, для соответствующего входного градиента можно передать None.Контекст можно использовать для получения тензоров, сохраненных во время прямого прохода. Он также содержит атрибут
ctx.needs_input_grad— кортеж логических значений, указывающих, для каких входных значений требуются градиенты. Например,backward()будет содержатьctx.needs_input_grad[0] = True, если для первого входного значенияforward()требуется вычислить градиент по отношению к выходному значению.- Тип возвращаемого значения:
-
static vmap(info, in_dims, *args)[исходный код] -
Определите поведение autograd.Function при использовании внутри
torch.vmap().Чтобы поддержать
torch.vmap()вtorch.autograd.Function(), необходимо либо переопределить этот статический метод, либо установитьgenerate_vmap_ruleвTrue(использовать оба варианта нельзя).Если вы решили переопределить этот статический метод, он должен принимать:
- объект
infoв качестве первого аргумента.info.batch_sizeзадает размерность оси, по которой выполняется vmap, аinfo.randomness— параметр случайности, переданный вtorch.vmap(). - кортеж
in_dimsв качестве второго аргумента. Для каждого аргумента вargsвin_dimsесть соответствующийOptional[int]. Его значение равноNone, если аргумент не является тензором или для него не выполняется vmap; в противном случае это целое число, задающее ось тензора, по которой выполняется vmap. -
*args, совпадающий с аргументамиforward().
Статический метод vmap возвращает кортеж из
(output, out_dims). Как иin_dims,out_dimsдолжен иметь ту же структуру, что иoutput, и содержать по одномуout_dimдля каждого выходного значения, указывая, есть ли у него измерение vmap и по какому индексу оно расположено.Подробнее см. в разделе Расширение torch.func с помощью autograd.Function.
- объект
-
© 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.NestedIOFunction.html