Вложенные разрывы графа
Создано: 28 июля 2025 г. | Последнее обновление: 3 декабря 2025 г.
Краткое описание:
- Разрывы графа во вложенных функциях могут приводить к труднопонимаемому поведению компилятора, которое описано ниже
- Вложенный разрыв графа приводит к дублированию разрыва графа
Напомним, что при применении torch.compile к функции трассируются также все вложенные вызовы функций. Вложенным разрывом графа называется любой разрыв графа, возникающий во вложенном вызове функции.
def inner(x):
...
torch._dynamo.graph_break() # nested graph break
...
@torch.compile
def outer(x):
...
y = inner(x)
...
Семантика возобновления после вложенных разрывов графа может быть непонятной, поэтому здесь мы описываем это поведение.
Напомним, что в fullgraph=False разрывы графа обрабатываются компиляцией сформированного на данный момент графа FX, выполнением неподдерживаемого кода на обычном Python, а затем возобновлением трассировки после неподдерживаемого кода с новым графом FX. Возобновление выполнения функции — довольно сложная техническая задача, поэтому возобновлять трассировку можно только для функций верхнего уровня.
Таким образом, с учетом этого ограничения мы можем возобновить трассировку после вложенного разрыва графа следующим образом:
Сначала рассмотрим приведенный ниже пример, в котором torch.compile выполняет трассировку, начиная с f, и продолжает ее до обнаружения разрыва графа в inner1.
def inner1(x):
x = x + 1
torch._dynamo.graph_break() # stop tracing due to graph break
return x + 2
def inner2(x):
x = x + 4
x = inner1(x)
x = x + 8
@torch.compile
def f(x):
# start tracing from here
x = x + 16
x = inner2(x)
x = x + 32
f(torch.randn(3))
Поскольку возобновить трассировку можно только из функций верхнего уровня, мы прерываем граф на вызове inner2 в f.
# The semantics of torch.compile(f)(x) is roughly this:
def compiled_f_semantics(x):
y = x + 16
z = inner2(y)
return torch.compile(resume_f_semantics)(z)
def resume_f_semantics(x):
return x + 32
compiled_f_semantics(torch.randn(3))
Затем inner2 автоматически компилируется как функция верхнего уровня. Мы продолжаем трассировку до повторного обнаружения разрыва графа в inner1.
def inner1(x):
x = x + 1
torch._dynamo.graph_break() # stop tracing due to graph break
return x + 2
# this torch.compile is automatically applied
@torch.compile
def inner2(x):
# start tracing from here
x = x + 4
x = inner1(x)
x = x + 8
def compiled_f_semantics(x):
y = x + 16
z = inner2(y)
return torch.compile(resume_f_semantics)(z)
def resume_f_semantics(x):
return x + 32
compiled_f_semantics(torch.randn(3))
Затем мы прерываем граф на вызове inner1 в inner2.
def compiled_inner2_semantics(x):
y = x + 4
z = inner1(y)
return torch.compile(resume_inner2_semantics)(z)
def resume_inner2_semantics(x):
return x + 8
Затем inner1 автоматически компилируется как функция верхнего уровня. Разрыв графа происходит в inner1, поэтому мы обрабатываем его обычным образом.
# this torch.compile is automatically applied
@torch.compile
def inner1(x):
# start tracing from here
x = x + 1
torch._dynamo.graph_break() # stop tracing due to graph break
return x + 2
def compiled_f_semantics(x):
y = x + 16
z = compiled_inner2_semantics(y)
return torch.compile(resume_f_semantics)(z)
def resume_f_semantics(x):
return x + 32
def compiled_inner2_semantics(x):
y = x + 4
z = inner1(y)
return torch.compile(resume_inner2_semantics)(z)
def resume_inner2_semantics(x):
return x + 8
compiled_f_semantics(torch.randn(3))
inner1 обрабатывается обычным образом:
def compiled_inner1_semantics(x):
y = x + 1
torch._dynamo.graph_break()
return torch.compile(resume_inner1_semantics)(y)
def resume_inner1_semantics(x):
return x + 2
Таким образом, исходный код семантически эквивалентен следующему:
def compiled_f_semantics(x):
y = x + 16
z = compiled_inner2_semantics(y)
return torch.compile(resume_f_semantics)(z)
def resume_f_semantics(x):
return x + 32
def compiled_inner2_semantics(x):
y = x + 4
z = compiled_inner1_semantics(y)
return torch.compile(resume_inner2_semantics)(z)
def resume_inner2_semantics(x):
return x + 8
def compiled_inner1_semantics(x):
y = x + 1
torch._dynamo.graph_break()
return torch.compile(resume_inner1_semantics)(y)
def resume_inner1_semantics(x):
return x + 2
compiled_f_semantics(torch.randn(3))
Обратите внимание: мы трассировали 3 функции верхнего уровня и трижды трассировали один и тот же разрыв графа. Это объясняет, почему при использовании torch.compile могут возникать дублирующиеся разрывы графа.
Итак, вложенные разрывы графа обрабатываются следующим образом:
- Трассировка от функции верхнего уровня до вложенного разрыва графа
- Разрыв графа в функции верхнего уровня при вызове функции второго уровня
- Компиляция отслеженных на данный момент операций PyTorch и выполнение скомпилированного графа
- Вызов функции второго уровня, которая автоматически компилируется как функция верхнего уровня
- Возобновление трассировки после вызова функции второго уровня
Обратите внимание, что время обработки этого разрыва графа составляет , где — глубина вложенности, а — число инструкций от функции верхнего уровня до разрыва графа. В итоге мы трассируем кадров и трассируем один и тот же разрыв графа раз.
© 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.nested_graph_breaks.html