Spec-Zone.ru › PyTorch 2.14

torch.compiler.nonstrict_trace

torch.compiler.nonstrict_trace(traceable_fn) [исходный код]

Декоратор для пометки функции как трассируемой в режиме nonstrict для dynamo.

Функция, трассируемая в режиме nonstrict, представляется в графе dynamo как непрозрачный вызов. Dynamo не трассирует тело функции (отсюда «nonstrict»), но aot_autograd трассирует его.

Это похоже на allow_in_graph, но с расширенной поддержкой: - Пользовательских классов в качестве входных данных (их необходимо зарегистрировать с помощью pytree) - nn.Module в качестве входных аргументов (параметры и буферы отслеживаются для autograd) - Глобальных тензоров и тензоров из замыканий, рассматриваемых как константы (предполагается, что во время выполнения они не обновляются)

Примечание

  • При использовании backend="eager" исходная функция Python выполняется напрямую. При использовании backend="aot_eager" выполняется граф, трассированный aot_autograd. При использовании backend="inductor" трассированный граф компилируется с помощью inductor.
  • Обучение поддерживается: можно вызвать .backward() для выходных данных, и градиенты будут проходить через функцию, трассируемую в режиме nonstrict.
Опасные шаблоны (могут привести к незаметным ошибкам):
  • Побочные эффекты между функцией, вызванной через nonstric_trace, и скомпилированной областью: функция не должна зависеть от переменных, изменяемых другим кодом внутри скомпилированной функции, а код после вызова не должен зависеть от внесённых ею изменений.
  • Неявные входные данные (замыкания/глобальные переменные): тензоры, захваченные из внешних областей видимости, рассматриваются как константы. Градиенты НЕ будут передаваться обратно к ним. Если нужны градиенты, передавайте тензоры в качестве явных аргументов.
Ограничения:
  • И входные, и выходные данные должны иметь типы, совместимые с pytree. Пользовательские классы необходимо зарегистрировать с помощью torch.utils._pytree.register_pytree_node(), torch.utils._pytree.register_dataclass() или torch.utils._pytree.register_constant(). Тензоры, примитивы Python (int, float, bool, str), символьные типы (SymInt, SymFloat, SymBool) и встроенные контейнеры (list, tuple, dict) уже поддерживаются по умолчанию.
  • Примитивные значения и структура контейнеров специализируются для каждой точки вызова: при каждом выполнении в этой точке вызова должны использоваться одни и те же примитивы и структура.

Пример:

>>> import torch
>>> @torch.compiler.nonstrict_trace
... def traced_forward(model, x):
...     # It's OK to have dynamo graph break within nonstrict_trace region
...     torch._dynamo.graph_break()
...     return model(x) + x
...
>>> class MyModule(torch.nn.Module):
...     def __init__(self):
...         super().__init__()
...         self.inner = torch.nn.Linear(10, 10)
...
...     def forward(self, x):
...         return traced_forward(self.inner, x)
...
>>> # Compile and run
>>> model = MyModule()
>>> opt_model = torch.compile(model, backend="aot_eager", fullgraph=True)
>>> out = opt_model(torch.randn(10, 10))
>>> out.sum().backward()  # Gradients flow through traced_forward
Тип возвращаемого значения:

Callable[[~_P], _R]

© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/generated/torch.compiler.nonstrict_trace.html

Spec-Zone.ru

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