Spec-Zone.ru › PyTorch 2.14

torch.export

Создано: 12 июня 2025 г. | Последнее обновление: 14 февраля 2026 г.

Обзор

torch.export.export() принимает torch.nn.Module и создает трассированный граф, представляющий только вычисления с тензорами в функции, в режиме опережающей компиляции (AOT); впоследствии его можно выполнить с другими входными данными или сериализовать.

import torch
from torch.export import export, ExportedProgram

class Mod(torch.nn.Module):
    def forward(self, x: torch.Tensor, y: torch.Tensor) -> torch.Tensor:
        a = torch.sin(x)
        b = torch.cos(y)
        return a + b

example_args = (torch.randn(10, 10), torch.randn(10, 10))

exported_program: ExportedProgram = export(Mod(), args=example_args)
print(exported_program)

torch.export создает чистое промежуточное представление (IR) со следующими инвариантами. Дополнительные спецификации IR приведены здесь.

  • Корректность: гарантируется, что это корректное представление исходной программы, сохраняющее соглашения о вызовах исходной программы.
  • Нормализованность: граф не содержит семантики Python. Подмодули исходных программ встраиваются в граф, образуя единый полностью развернутый вычислительный граф.
  • Свойства графа: по умолчанию граф может содержать как функциональные, так и нефункциональные операторы (включая операторы с изменением данных). Чтобы получить чисто функциональный граф, используйте run_decompositions(), который удаляет изменения и псевдонимы.
  • Метаданные: граф содержит метаданные, полученные во время трассировки, например трассировку стека из кода пользователя.

Внутри torch.export использует следующие новейшие технологии:

  • TorchDynamo (torch._dynamo) — это внутренний API, использующий возможность CPython под названием Frame Evaluation API для безопасной трассировки графов PyTorch. Это значительно улучшает процесс захвата графа и сокращает количество изменений, необходимых для полной трассировки кода PyTorch.
  • AOT Autograd гарантирует декомпозицию/понижение графа до набора операторов ATen. При использовании run_decompositions() также может выполняться функционализация.
  • Torch FX (torch.fx) — базовое представление графа, позволяющее выполнять гибкие преобразования на Python.

Существующие фреймворки

torch.compile() также использует тот же стек PT2, что и torch.export, но имеет некоторые отличия:

  • JIT и AOT: torch.compile() — это JIT-компилятор, который не предназначен для создания скомпилированных артефактов вне процесса развертывания.
  • Частичный и полный захват графа: если torch.compile() сталкивается с частью модели, которую невозможно трассировать, происходит «разрыв графа», и программа продолжает выполняться в обычной среде выполнения Python. В отличие от него, torch.export стремится получить полное представление графа модели PyTorch и поэтому завершится с ошибкой при обнаружении нетрассируемого элемента. Поскольку torch.export создает полный граф, не зависящий от каких-либо возможностей или среды выполнения Python, этот граф можно сохранять, загружать и запускать в разных средах и на разных языках.
  • Компромисс в удобстве использования: поскольку torch.compile() может переключаться на среду выполнения Python при обнаружении нетрассируемого элемента, он гораздо гибче. В отличие от него, torch.export требует от пользователей предоставить дополнительную информацию или переписать код, чтобы сделать его трассируемым.

В отличие от torch.fx.symbolic_trace(), torch.export выполняет трассировку с помощью TorchDynamo, который работает на уровне байт-кода Python и поэтому способен трассировать произвольные конструкции Python, не ограничиваясь возможностями перегрузки операторов Python. Кроме того, torch.export тщательно отслеживает метаданные тензоров, поэтому условные выражения, зависящие, например, от формы тензора, не приводят к сбою трассировки. В целом ожидается, что torch.export будет работать с большим количеством пользовательских программ и создавать графы более низкого уровня (на уровне операторов torch.ops.aten). Обратите внимание, что пользователи могут по-прежнему использовать torch.fx.symbolic_trace() в качестве этапа предварительной обработки перед torch.export.

В отличие от torch.jit.script(), torch.export не захватывает поток управления или структуры данных Python, если не используются явные операторы управления потоком, но поддерживает больше возможностей языка Python благодаря всестороннему охвату байт-кодов Python. Полученные графы проще и содержат только линейный поток управления, за исключением явных операторов управления потоком.

В отличие от torch.jit.trace(), torch.export является корректным: он может трассировать код, выполняющий целочисленные вычисления размеров, и фиксирует все условия, необходимые для того, чтобы конкретная трассировка оставалась допустимой для других входных данных.

Экспорт модели PyTorch

Основная точка входа — torch.export.export(). Она принимает torch.nn.Module и примеры входных данных и захватывает вычислительный граф в torch.export.ExportedProgram. Пример:

import torch
from torch.export import export, ExportedProgram

# Simple module for demonstration
class M(torch.nn.Module):
    def __init__(self) -> None:
        super().__init__()
        self.conv = torch.nn.Conv2d(
            in_channels=3, out_channels=16, kernel_size=3, padding=1
        )
        self.relu = torch.nn.ReLU()
        self.maxpool = torch.nn.MaxPool2d(kernel_size=3)

    def forward(self, x: torch.Tensor, *, constant=None) -> torch.Tensor:
        a = self.conv(x)
        a.add_(constant)
        return self.maxpool(self.relu(a))

example_args = (torch.randn(1, 3, 256, 256),)
example_kwargs = {"constant": torch.ones(1, 16, 256, 256)}

exported_program: ExportedProgram = export(
    M(), args=example_args, kwargs=example_kwargs
)
print(exported_program)

# To run the exported program, we can use the `module()` method
print(exported_program.module()(torch.randn(1, 3, 256, 256), constant=torch.ones(1, 16, 256, 256)))

Изучив ExportedProgram, можно отметить следующее:

  • torch.fx.Graph содержит вычислительный граф исходной программы и записи исходного кода для удобства отладки.
  • Граф содержит только операторы torch.ops.aten, перечисленные здесь, а также пользовательские операторы.
  • Параметры (weight и bias свертки) переносятся во входные данные графа, поэтому в графе нет узлов get_attr, которые ранее присутствовали в результате torch.fx.symbolic_trace().
  • torch.export.ExportGraphSignature описывает сигнатуры входных и выходных данных, а также указывает, какие входные данные являются параметрами.
  • Для каждого узла графа указаны форма и тип данных создаваемых тензоров. Например, узел conv2d создает тензор с типом данных torch.float32 и формой (1, 16, 256, 256).

Задание динамических размеров

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

Пример:

import torch
import traceback as tb

class M(torch.nn.Module):
    def __init__(self):
        super().__init__()

        self.branch1 = torch.nn.Sequential(
            torch.nn.Linear(64, 32), torch.nn.ReLU()
        )
        self.branch2 = torch.nn.Sequential(
            torch.nn.Linear(128, 64), torch.nn.ReLU()
        )
        self.buffer = torch.ones(32)

    def forward(self, x1, x2):
        out1 = self.branch1(x1)
        out2 = self.branch2(x2)
        return (out1 + self.buffer, out2)

example_args = (torch.randn(32, 64), torch.randn(32, 128))

ep = torch.export.export(M(), example_args)
print(ep)

example_args2 = (torch.randn(64, 64), torch.randn(64, 128))
try:
    ep.module()(*example_args2)  # fails
except Exception:
    tb.print_exc()

Однако некоторые измерения, например размер пакета, могут быть динамическими и меняться от запуска к запуску. Такие измерения необходимо задать с помощью API torch.export.Dim(), создав их и передав в torch.export.export() через аргумент dynamic_shapes.

import torch

class M(torch.nn.Module):
    def __init__(self):
        super().__init__()

        self.branch1 = torch.nn.Sequential(
            torch.nn.Linear(64, 32), torch.nn.ReLU()
        )
        self.branch2 = torch.nn.Sequential(
            torch.nn.Linear(128, 64), torch.nn.ReLU()
        )
        self.buffer = torch.ones(32)

    def forward(self, x1, x2):
        out1 = self.branch1(x1)
        out2 = self.branch2(x2)
        return (out1 + self.buffer, out2)

example_args = (torch.randn(32, 64), torch.randn(32, 128))

# Create a dynamic batch size
batch = torch.export.Dim("batch")
# Specify that the first dimension of each input is that batch size
dynamic_shapes = {"x1": {0: batch}, "x2": {0: batch}}

ep = torch.export.export(
    M(), args=example_args, dynamic_shapes=dynamic_shapes
)
print(ep)

example_args2 = (torch.randn(64, 64), torch.randn(64, 128))
ep.module()(*example_args2)  # success

Следует также отметить следующее:

  • С помощью API torch.export.Dim() и аргумента dynamic_shapes мы указали, что первое измерение каждого входного тензора является динамическим. Входные данные x1 и x2 имеют символьные формы (s0, 64) и (s0, 128) вместо форм (32, 64) и (32, 128), заданных для входных примеров. s0 — это символ, обозначающий, что данное измерение может принимать диапазон значений.
  • exported_program.range_constraints описывает диапазоны каждого символа, встречающегося в графе. В данном случае s0 имеет диапазон [0, int_oo]. По техническим причинам, которые сложно объяснить здесь, предполагается, что значения не равны 0 или 1. Это не ошибка и не обязательно означает, что экспортированная программа не будет работать при размерах 0 или 1. Подробное обсуждение этой темы см. в разделе Проблема специализации для 0/1.

В этом примере для создания динамического измерения мы использовали Dim("batch"). Это самый явный способ задать динамичность. Также можно использовать Dim.DYNAMIC и Dim.AUTO. Оба метода рассмотрены в следующем разделе.

Именованные измерения

Для каждого измерения, заданного с помощью Dim("name"), будет выделена символьная форма. Если задать Dim с тем же именем, будет сгенерирован тот же символ. Это позволяет пользователям указывать, какие символы выделяются для каждого измерения входных данных.

batch = Dim("batch")
dynamic_shapes = {"x1": {0: dim}, "x2": {0: batch}}

Для каждого Dim можно задать минимальное и максимальное значения. Также можно задавать отношения между Dim в линейных выражениях от одной переменной: A * dim + B. Это позволяет задавать более сложные ограничения для динамических измерений, например делимость на целое число. Благодаря этим возможностям пользователи могут явно ограничивать динамическое поведение создаваемого ExportedProgram.

dx = Dim("dx", min=4, max=256)
dh = Dim("dh", max=512)
dynamic_shapes = {
    "x": (dx, None),
    "y": (2 * dx, dh),
}

Однако ConstraintViolationErrors будет вызвана, если во время трассировки появятся условия, противоречащие заданным отношениям или спецификациям статических/динамических размеров. Например, в приведенной выше спецификации утверждается следующее:

  • x.shape[0] имеет диапазон [4, 256] и связано с y.shape[0] отношением y.shape[0] == 2 * x.shape[0].
  • x.shape[1] является статическим.
  • y.shape[1] имеет диапазон [0, 512] и не связано ни с каким другим измерением.

Если во время трассировки обнаружится, что какое-либо из этих утверждений неверно (например, x.shape[0] является статическим, или y.shape[1] имеет меньший диапазон, или y.shape[0] != 2 * x.shape[0]), будет вызвано исключение ConstraintViolationError, и пользователю потребуется изменить спецификацию dynamic_shapes.

Подсказки для измерений

Вместо явного задания динамичности с помощью Dim("name") можно позволить torch.export выводить диапазоны и отношения динамических значений с помощью Dim.DYNAMIC. Это также более удобный способ задать динамичность, если вы не знаете точно, насколько динамичны ваши значения.

dynamic_shapes = {
    "x": (Dim.DYNAMIC, None),
    "y": (Dim.DYNAMIC, Dim.DYNAMIC),
}

Для Dim.DYNAMIC можно также задать минимальные и максимальные значения, которые будут использоваться как подсказки для экспорта. Однако если во время трассировки экспорта будет обнаружен другой диапазон, он автоматически обновится без вызова ошибки. Задавать отношения между динамическими значениями нельзя. Вместо этого экспорт выведет их самостоятельно, а пользователи смогут увидеть их, изучив утверждения в графе. При таком способе задания динамичности исключение ConstraintViolationErrors будет вызвано только в том случае, если указанное значение будет распознано как статическое.

Еще удобнее задавать динамичность с помощью Dim.AUTO. Он работает как Dim.DYNAMIC, но не вызывает ошибку, если измерение распознано как статическое. Это полезно, если вы не знаете, какие именно значения являются динамическими, и хотите экспортировать программу, применяя динамический подход по принципу «по возможности».

ShapesCollection

При указании динамических входных данных с помощью dynamic_shapes необходимо задать динамичность каждого входного значения. Например, для следующих входных данных:

args = {"x": tensor_x, "others": [tensor_y, tensor_z]}

нам потребуется задать динамичность tensor_x, tensor_y и tensor_z, а также динамические формы:

# With named-Dims
dim = torch.export.Dim(...)
dynamic_shapes = {"x": {0: dim, 1: dim + 1}, "others": [{0: dim * 2}, None]}

torch.export(..., args, dynamic_shapes=dynamic_shapes)

Однако это довольно сложно, поскольку спецификацию dynamic_shapes нужно задавать в той же вложенной структуре, что и входные аргументы. Вместо этого проще задавать динамические формы с помощью вспомогательной утилиты torch.export.ShapesCollection: она позволяет не указывать динамичность каждого входного значения, а непосредственно назначать динамические измерения для входных данных.

import torch

class M(torch.nn.Module):
    def forward(self, inp):
        x = inp["x"] * 1
        y = inp["others"][0] * 2
        z = inp["others"][1] * 3
        return x, y, z

tensor_x = torch.randn(3, 4, 8)
tensor_y = torch.randn(6)
tensor_z = torch.randn(6)
args = {"x": tensor_x, "others": [tensor_y, tensor_z]}

dim = torch.export.Dim("dim")
sc = torch.export.ShapesCollection()
sc[tensor_x] = (dim, dim + 1, 8)
sc[tensor_y] = {0: dim * 2}

print(sc.dynamic_shapes(M(), (args,)))
ep = torch.export.export(M(), (args,), dynamic_shapes=sc)
print(ep)

AdditionalInputs

Если вы не знаете, насколько динамичны входные данные, но располагаете достаточным набором тестовых данных или данных профилирования, позволяющим определить характерные входные данные модели, вместо dynamic_shapes можно использовать torch.export.AdditionalInputs. Можно указать все возможные входные данные, используемые для трассировки программы, и AdditionalInputs определит, какие из них являются динамическими, на основе изменений форм входных данных.

Пример:

import dataclasses
import torch
import torch.utils._pytree as pytree

@dataclasses.dataclass
class D:
    b: bool
    i: int
    f: float
    t: torch.Tensor

pytree.register_dataclass(D)

class M(torch.nn.Module):
    def forward(self, d: D):
        return d.i + d.f + d.t

input1 = (D(True, 3, 3.0, torch.ones(3)),)
input2 = (D(True, 4, 3.0, torch.ones(4)),)
ai = torch.export.AdditionalInputs()
ai.add(input1)
ai.add(input2)

print(ai.dynamic_shapes(M(), input1))
ep = torch.export.export(M(), input1, dynamic_shapes=ai)
print(ep)

Сериализация

Чтобы сохранить ExportedProgram, пользователи могут использовать API torch.export.save() и torch.export.load(). В результате создается ZIP-архив с определенной структурой. Подробности о структуре приведены в спецификации архива PT2.

Пример:

import torch

class MyModule(torch.nn.Module):
    def forward(self, x):
        return x + 10

exported_program = torch.export.export(MyModule(), (torch.randn(5),))

torch.export.save(exported_program, 'exported_program.pt2')
saved_exported_program = torch.export.load('exported_program.pt2')

IR экспорта: обучение и инференс

Граф, создаваемый torch.export, содержит только операторы ATen — базовые единицы вычислений в PyTorch. В зависимости от сценария использования Export предоставляет разные уровни IR:

Тип IR

Как получить

Свойства

Количество операторов

Сценарий использования

IR обучения

torch.export.export() (по умолчанию)

Может содержать изменения

~3000

Обучение с autograd

IR инференса

ep.run_decompositions(decomp_table={})

Чисто функциональный

~2000

Развертывание инференса

Core ATen IR

ep.run_decompositions(decomp_table=None)

Чисто функциональный, с высокой степенью декомпозиции

~180

Минимальная поддержка бэкенда

IR обучения (по умолчанию)

По умолчанию экспорт создает IR обучения, содержащий все операторы ATen, в том числе функциональные и нефункциональные (изменяющие данные). Функциональный оператор не изменяет входные данные и не создает для них псевдонимы, тогда как нефункциональные операторы могут изменять входные данные на месте. Список всех операторов ATen можно найти здесь, а функциональность оператора можно проверить с помощью op._schema.is_mutable.

Этот IR обучения может содержать изменения, предназначен для сценариев обучения и может использоваться с PyTorch Autograd в обычном режиме выполнения.

import torch

class M(torch.nn.Module):
    def __init__(self) -> None:
        super().__init__()
        self.conv = torch.nn.Conv2d(1, 3, 1, 1)
        self.bn = torch.nn.BatchNorm2d(3)

    def forward(self, x):
        x = self.conv(x)
        x = self.bn(x)
        return (x,)

ep_for_training = torch.export.export(M(), (torch.randn(1, 1, 3, 3),))
print(ep_for_training.graph_module.print_readable(print_output=False))

IR инференса (с помощью run_decompositions)

Чтобы получить IR инференса, пригодный для развертывания, используйте API ExportedProgram.run_decompositions(). Этот метод автоматически выполняет следующие действия:

  1. Функционализирует граф (удаляет все изменения и заменяет их функциональными эквивалентами).
  2. При необходимости выполняет декомпозицию операторов ATen на основе предоставленной таблицы декомпозиций.

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

Если указать пустую таблицу декомпозиций (decomp_table={}), будет выполнена только функционализация без дополнительной декомпозиции. В результате получится IR инференса примерно с 2000 функциональными операторами (по сравнению с более чем 3000 в IR обучения).

import torch

class M(torch.nn.Module):
    def __init__(self) -> None:
        super().__init__()
        self.conv = torch.nn.Conv2d(1, 3, 1, 1)
        self.bn = torch.nn.BatchNorm2d(3)

    def forward(self, x):
        x = self.conv(x)
        x = self.bn(x)
        return (x,)

ep_for_training = torch.export.export(M(), (torch.randn(1, 1, 3, 3),))
with torch.no_grad():
    ep_for_inference = ep_for_training.run_decompositions(decomp_table={})
print(ep_for_inference.graph_module.print_readable(print_output=False))

Как видно, ранее выполнявшийся на месте оператор torch.ops.aten.add_.default заменен функциональным оператором torch.ops.aten.add.default.

Core ATen IR

Далее IR инференса можно понизить до Core ATen Operator Set <https://docs.pytorch.org/docs/main/user_guide/torch_compiler/torch.compiler_ir.html#core-aten-ir>__, содержащего всего около 180 операторов. Для этого нужно передать decomp_table=None (использующую таблицу декомпозиций по умолчанию) в run_decompositions(). Этот IR оптимален для бэкендов, стремящихся свести к минимуму количество операторов, которые необходимо реализовать.

import torch

class M(torch.nn.Module):
    def __init__(self) -> None:
        super().__init__()
        self.conv = torch.nn.Conv2d(1, 3, 1, 1)
        self.bn = torch.nn.BatchNorm2d(3)

    def forward(self, x):
        x = self.conv(x)
        x = self.bn(x)
        return (x,)

ep_for_training = torch.export.export(M(), (torch.randn(1, 1, 3, 3),))
with torch.no_grad():
    core_aten_ir = ep_for_training.run_decompositions(decomp_table=None)
print(core_aten_ir.graph_module.print_readable(print_output=False))

Теперь видно, что torch.ops.aten.conv2d.default декомпозирован в torch.ops.aten.convolution.default. Это объясняется тем, что convolution — более «базовый» оператор: такие операции, как conv1d и conv2d, можно реализовать с помощью одного и того же оператора.

Также можно задать собственные правила декомпозиции:

class M(torch.nn.Module):
    def __init__(self) -> None:
        super().__init__()
        self.conv = torch.nn.Conv2d(1, 3, 1, 1)
        self.bn = torch.nn.BatchNorm2d(3)

    def forward(self, x):
        x = self.conv(x)
        x = self.bn(x)
        return (x,)

ep_for_training = torch.export.export(M(), (torch.randn(1, 1, 3, 3),))

my_decomp_table = torch.export.default_decompositions()

def my_awesome_custom_conv2d_function(x, weight, bias, stride=[1, 1], padding=[0, 0], dilation=[1, 1], groups=1):
    return 2 * torch.ops.aten.convolution(x, weight, bias, stride, padding, dilation, False, [0, 0], groups)

my_decomp_table[torch.ops.aten.conv2d.default] = my_awesome_custom_conv2d_function
my_ep = ep_for_training.run_decompositions(my_decomp_table)
print(my_ep.graph_module.print_readable(print_output=False))

Обратите внимание: вместо декомпозиции torch.ops.aten.conv2d.default в torch.ops.aten.convolution.default теперь выполняется декомпозиция в torch.ops.aten.convolution.default и torch.ops.aten.mul.Tensor в соответствии с нашим пользовательским правилом.

Ограничения torch.export

Поскольку torch.export за один проход захватывает вычислительный граф программы PyTorch, при трассировке могут обнаружиться нетрассируемые части программы: практически невозможно обеспечить поддержку трассировки всех возможностей PyTorch и Python. В случае torch.compile неподдерживаемая операция приведет к «разрыву графа», после чего она будет выполнена с помощью стандартного механизма вычисления Python. В отличие от него, torch.export требует от пользователей предоставить дополнительную информацию или переписать части кода, чтобы сделать их трассируемыми.

Draft-export — полезный ресурс для выявления разрывов графа, которые возникнут при трассировке программы, а также для получения дополнительной отладочной информации, необходимой для устранения ошибок.

ExportDB также поможет узнать, какие типы программ поддерживаются, а какие нет, и как переписать программы, чтобы сделать их трассируемыми.

Неподдерживаемые возможности TorchDynamo

При использовании torch.export с strict=True для оценки программы на уровне байт-кода Python и ее трассировки в граф используется TorchDynamo. По сравнению с предыдущими средствами трассировки потребуется значительно меньше изменений, чтобы сделать программу трассируемой, однако некоторые возможности Python все еще не поддерживаются. Чтобы обойти разрывы графа, можно выполнить нестрогий экспорт, изменив флаг strict на strict=False.

Поток управления, зависящий от данных/формы

Разрывы графа также могут возникать при потоках управления, зависящих от данных (if x.shape[0] > 2), если формы не специализируются: компилятор трассировки не может обработать такие потоки, не генерируя код для комбинаторно растущего числа путей. В таких случаях пользователям потребуется переписать код с помощью специальных операторов управления потоком. Сейчас мы поддерживаем операторы высшего порядка для представления таких конструкций управления потоком, как условные выражения, отображение, сканирование и циклы.

Дополнительные способы устранения ошибок, зависящих от данных, описаны в этом руководстве.

Отсутствующие Fake/Meta-ядра операторов

Для всех операторов при трассировке требуется ядро FakeTensor (также называемое Meta-ядром). Оно используется для определения форм входных и выходных данных оператора.

Дополнительные сведения см. в этом руководстве.

Если в вашей модели используется оператор ATen, для которого пока не реализовано ядро FakeTensor, создайте сообщение о проблеме.

Дополнительные материалы

Дополнительные ссылки для пользователей Export

  • Справочник API torch.export
  • Модель программирования torch.export
  • Спецификация IR torch.export
  • Спецификация архива PT2
  • Черновой экспорт
  • Совместная работа с дескрипторами
  • Операторы управления потоком
  • ExportDB
  • AOTInductor: опережающая компиляция моделей, экспортированных с помощью Torch
  • IR

Подробное руководство для разработчиков PyTorch

  • Динамические формы
  • Тензор Fake
  • Создание преобразований графов для IR ATen

© 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/export.html

Spec-Zone.ru

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