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) уже поддерживаются по умолчанию. - Примитивные значения и структура контейнеров специализируются для каждой точки вызова: при каждом выполнении в этой точке вызова должны использоваться одни и те же примитивы и структура.
- И входные, и выходные данные должны иметь типы, совместимые с pytree. Пользовательские классы необходимо зарегистрировать с помощью
Пример:
>>> 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