Spec-Zone.ru › PyTorch 2.14

Модель программирования с нестрогой трассировкой

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

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

  • Нестрогая трассировка — способ трассировки кода Python, менее строгий, чем Dynamo, но способный привести к незаметным ошибкам.
  • При нестрогой трассировке выполняется функция Python, а возможности перегрузки операторов в Python и PyTorch используются для записи в трассу операций с тензорами, произошедших во время выполнения.
  • Функция является трассируемой с помощью нестрогой трассировки, если она соответствует некоторым ограничениям: в частности, функция должна быть чистой и не должна напрямую манипулировать Tensor.data_ptr().
  • При нестрогой трассировке для определённых переменных может выполняться специализация, и они могут рассматриваться как константы, при этом значения переменных встраиваются в трассу.

Внутренние компоненты torch.compile (make_fx, AOTDispatcher) используют нестрогую трассировку. torch._dynamo.nonstrict_trace также можно использовать в коде torch.compiled, чтобы помечать участки кода, которые нужно трассировать с помощью нестрогой трассировки. При нестрогой трассировке выполняется функция Python, а возможности перегрузки операторов в Python и PyTorch используются для записи в трассу операций с тензорами, произошедших во время выполнения.

make_fx — основная точка входа для нестрогой трассировки. При выполнении следующей функции для заданных входных данных выбирается только верхняя ветвь, поэтому захватывается граф, содержащий только эту ветвь.

from torch.fx.experimental.proxy_tensor import make_fx
def f(x):
    if x.shape[0] > 2:
        return x ** 2 / 6
    else:
        return x * 3
x = torch.randn(3)
gm = make_fx(f, tracing_mode="fake")(x)
gm.print_readable()

Нестрогая трассировка отличается от трассировки Dynamo (строгой трассировки) тем, что она небезопасна: для заданной функции она захватывает граф операций с тензорами, семантика которого может отличаться от семантики исходной функции. Для заданной функции Python трассировка Dynamo захватывает граф операций с тензорами и остаточный байт-код, которые вместе обеспечивают ту же семантику, что и функция Python.

Чистые функции

Нестрогая трассировка корректна только для чистых функций, поэтому с помощью нестрогой трассировки следует трассировать только чистые функции.

Чистая функция обладает следующими свойствами:

  • Детерминированность. При одинаковых входных данных чистая функция всегда возвращает одинаковый результат.
  • Отсутствие побочных эффектов. Чистая функция не имеет побочных эффектов, таких как изменение внешнего состояния или выполнение операций ввода-вывода.
  • Явные входные и выходные данные. Все входные данные должны передаваться через параметры функции, а все выходные данные должны возвращаться из функции.

Ниже приведены примеры нечистых функций, для которых захваченный граф ведёт себя иначе, чем исходная функция.

Пример 1: Нет явных входных данных (например, обращение к глобальному тензору)

var = torch.tensor(1)
def function_with_global_access(y):
    return y + var
x = torch.tensor([0, 1, 2])
# _allow_non_fake_inputs=True is needed to capture the global variable
# for demonstration purposes.
gm = make_fx(
    function_with_global_access, tracing_mode="fake", _allow_non_fake_inputs=True
)(x)
# Non-strict Tracing captures the value of the global (1.)
print("1. call function", function_with_global_access(x))
print("1. call graph", gm(x))
# However, after changing the global, the captured graph
# produces a different result from the original function
var = torch.tensor(2)
print("2. call function", function_with_global_access(x))
print("2. call graph", gm(x))
# To capture a graph that can have a varying `var` tensor,
# it must be an explicit input:
def function_fixed(y, var):
    return y + var
var = torch.tensor(3)
gm = make_fx(function_fixed, tracing_mode="fake")(x, var)
print("3. call function", function_fixed(x, var))
print("3. call graph", gm(x, var))
var = torch.tensor(4)
print("4. call function", function_fixed(x, var))
print("4. call graph", gm(x, var))

Объяснение приведено в разделе Специализация и константы.

Пример 2: Побочный эффект (вывод на экран)

def function_with_side_effect(y):
    print(y)
x = torch.tensor([0, 1, 2])
_ = function_with_side_effect(x)

Выполнение f в Python выводит тензор на экран в качестве побочного эффекта.

gm = make_fx(function_with_side_effect, tracing_mode="fake")(x)

При нестрогой трассировке этот вывод выполняется во время захвата графа.

_ = gm(x)

Граф не сохраняет вызов инструкции print, поэтому при выполнении графа ничего не выводится.

Пример 3: Побочный эффект (изменение входного списка)

lst = []
def function_with_input_list_mutation(lst):
    val = lst.pop()
    return val
x = torch.tensor([0, 1, 2])
y = torch.tensor([0, 1, 2])
# Each time the function is executed, the list shrinks in size
lst = [x, y]
function_with_input_list_mutation(lst)
print("len(lst) after one call", len(lst))
function_with_input_list_mutation(lst)
print("len(lst) after two calls", len(lst))
# With Non-strict Tracing, the length of the list shrinks during
# the graph capture but not in invocations of the graph.
lst = [x, y]
gm = make_fx(function_with_input_list_mutation, tracing_mode="fake")(lst)
print("len(lst) after graph capture", len(lst))
gm(lst)
print("len(lst) after one call to graph", len(lst))
gm(lst)
print("len(lst) after two calls to graph", len(lst))

Отсутствие прямого обращения к data_ptr

Прямое манипулирование Tensor.data_ptr не поддерживается нестрогой трассировкой. Причина в том, что PyTorch не может определить, как именно вы манипулировали data_ptr.

import ctypes
# Create a tensor with a single element
tensor = torch.tensor([42], dtype=torch.int32)  # Using int32 for simplicity
def function_with_data_ptr(tensor):
    # Get the data pointer
    ptr = tensor.data_ptr()
    # Cast the pointer to a ctypes pointer
    ctypes_ptr = ctypes.cast(ptr, ctypes.POINTER(ctypes.c_int32))
    # Increment the value at the pointer
    ctypes_ptr.contents.value += 1
    return tensor
try:
    make_fx(function_with_data_ptr, tracing_mode="fake")(tensor)
except Exception as e:
    print(e)

Специализация и константы

При нестрогой трассировке захватывается граф, который может быть специализирован для некоторых значений. Это означает, что захваченный граф действителен только для этих значений. Говорят, что граф рассматривает эти значения как константы.

При нестрогой трассировке все переменные, не являющиеся тензорами, считаются константами:

def f(x, y):
    return x + y
x = torch.tensor([0, 1, 2])
y = 3.14
gm = make_fx(f, tracing_mode="fake")(x, y)
gm.print_readable()

3.14 — константа в графе.

Нестрогая трассировка также выполняет специализацию по свойствам входных тензоров.

def f(x):
    if x.shape[0] > 2:
        return x ** 2 / 6
    else:
        return x * 3
x = torch.randn(3)
gm = make_fx(f, tracing_mode="fake")(x)
gm.print_readable()

Кроме того, специализация выполняется по любым переменным, которые не передаются в функцию напрямую:

var = torch.tensor(1)
def f(x):
    return x + y
x = torch.randn(3)
gm = make_fx(f, tracing_mode="fake")(x)
gm.print_readable()

© 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.non_strict_tracing_model.html

Spec-Zone.ru

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