Spec-Zone.ru › PyTorch 2.14
fullgraph=False">

Вложенные разрывы графа

Создано: 28 июля 2025 г. | Последнее обновление: 3 декабря 2025 г.

Краткое описание:

  • Разрывы графа во вложенных функциях могут приводить к труднопонимаемому поведению компилятора, которое описано ниже
  • Вложенный разрыв графа приводит к O(N)\mathcal O(N) дублированию разрыва графа

Напомним, что при применении 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 и выполнение скомпилированного графа
  • Вызов функции второго уровня, которая автоматически компилируется как функция верхнего уровня
  • Возобновление трассировки после вызова функции второго уровня

Обратите внимание, что время обработки этого разрыва графа составляет O(NK)\mathcal O(NK), где NN — глубина вложенности, а KK — число инструкций от функции верхнего уровня до разрыва графа. В итоге мы трассируем O(N2)\mathcal O(N^2) кадров и трассируем один и тот же разрыв графа O(N)\mathcal O(N) раз.

© 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

Spec-Zone.ru

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