Пропускаемые функции
Создано: 28 июл. 2025 | Последнее обновление: 03 дек. 2025
Краткое описание:
- Иногда
torch.compileполностью отказывается компилировать функцию и вместо этого выполняет её в обычном режиме, что может привести к упущенным возможностям оптимизации. - Существуют способы обойти пропуск функций и возобновить трассировку вокруг проблемного кода.
Иногда torch.compile с fullgraph=False не может возобновить трассировку при возникновении разрыва графа или другой ошибки компилятора. Во многих таких случаях torch.compile пропускает компиляцию всей функции и выполняет её в обычном режиме.
Обратите внимание: пропуск применяется только к текущей функции, а НЕ к вызовам вложенных функций. torch.compile всё равно попытается скомпилировать вложенные вызовы.
def inner1(x):
return x + 1
def inner2(x):
return x + 2
@torch.compile
def fn(x):
x = inner1(x)
torch._dynamo.skip_frame()
x = inner2(x)
fn(torch.randn(3))
В примере выше torch.compile трассирует fn (включая inner1) до skip_frame. Затем fn пропускается и выполняется в обычном режиме — inner1 и inner2 компилируются при вызове.
Пропуск функций может привести к упущенным возможностям оптимизации, поэтому важно проверять, не пропускается ли код, который нужно скомпилировать, и, если это так, обходить пропуск.
Разрыв графа в цикле
torch.compile не может возобновить трассировку, если разрыв графа происходит в цикле:
@torch.compile
def fn(x):
for i in range(5):
x = x + 1
if i == 3:
torch._dynamo.graph_break()
return x
fn(torch.randn(3))
В этом примере можно избежать пропуска, развернув цикл:
@torch.compile
def fn(x):
def inner(i):
nonlocal x
x = x + 1
if i == 3:
torch._dynamo.graph_break()
inner(0)
inner(1)
inner(2)
inner(3)
inner(4)
return x
fn(torch.randn(3))
Как правило, устранение разрыва графа, вызывающего пропуск, также устраняет и сам пропуск.
Разрыв графа в менеджере контекста
Ещё один распространённый пример разрыва графа, после которого невозможно возобновить трассировку, — разрыв графа в большинстве менеджеров контекста:
class CustomCtxManager:
def __enter__(self):
pass
def __exit__(self, exc_type, exc_value, traceback):
pass
@torch.compile
def fn(x):
with CustomCtxManager():
x = x + 1
torch._dynamo.graph_break()
return x + 1
fn(torch.randn(3))
Можно избежать пропуска, переместив разрыв графа за пределы менеджера контекста:
@torch.compile
def fn(x):
with CustomCtxManager():
x = x + 1
torch._dynamo.graph_break()
with CustomCtxManager():
return x + 1
fn(torch.randn(3))
Существуют менеджеры контекста, после разрыва графа в которых Dynamo может возобновить трассировку. Некоторые из них можно найти в supported_ctx_manager_classes в torch/_dynamo/variables/torch.py. В целом поддерживают возобновление после разрыва графа все менеджеры контекста, представленные подклассом ContextWrappingVariable в torch/_dynamo/variables/ctx_manager.py. Например:
import contextlib
@torch.compile
def fn(x):
with contextlib.nullcontext():
with torch.no_grad():
x = x + 1
torch._dynamo.graph_break()
return x + 1
fn(torch.randn(3))
Разрыв графа в блоке try
После разрыва графа в блоке try невозможно возобновить трассировку:
@torch.compile
def fn(x):
try:
x = x + 1
torch._dynamo.graph_break()
return x + 1
except Exception as e:
pass
fn(torch.randn(3))
Можно избежать пропуска, переместив разрыв графа за пределы блока try:
@torch.compile
def fn(x):
try:
x = x + 1
except Exception as e:
pass
torch._dynamo.graph_break()
try:
return x + 1
except Exception as e:
pass
fn(torch.randn(3))
Достижение предела перекомпиляции
Ошибки компилятора
Некоторые ошибки компилятора приводят к пропуску функций. Другие ошибки компилятора приводят к критической ошибке, а не к пропуску функции.
Работа с пропускаемыми функциями
Как правило, пропуск функции можно устранить, исправив исходный разрыв графа или ошибку, из-за которой функция пропускается.
Если разрыв графа или ошибку, вызывающие пропуск функции, трудно исправить, попробуйте изолировать их в отдельной функции, чтобы пропускалось как можно меньше кода.
def inner1(x):
return x + 1
def inner2(x):
return x + 2
@torch.compile
def fn(x):
x = inner1(x)
def problematic_code():
torch._dynamo.skip_frame()
problematic_code()
x = inner2(x)
fn(torch.randn(3))
© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/user_guide/torch_compiler/compile/programming_model.skipped_functions.html