Экспорт модели из PyTorch через torch.onnx.export

Экспорт модели PyTorch в формат ONNX является первым шагом для последующего запуска модели в браузере с использованием ONNX Runtime Web. Этот процесс предполагает преобразование динамической модели PyTorch в статический граф вычислений, понятный ONNX, с сохранением всех слоёв и параметров.

Подготовка модели

Перед экспортом важно убедиться, что модель переведена в режим inference с помощью метода .eval(). Это отключает поведение, характерное для обучения, такое как дропаут или нормализация батчей:

import torch
import torchvision.models as models

model = models.resnet18(pretrained=True)
model.eval()

Формирование входного тензора

ONNX требует конкретного размера входного тензора для корректного построения графа. Обычно используется пример батча фиксированного размера:

dummy_input = torch.randn(1, 3, 224, 224)  # 1 изображение, 3 канала, 224x224

Важно, чтобы размерность соответствовала ожидаемой модели.

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

Метод torch.onnx.export выполняет сериализацию модели в файл .onnx. Основные параметры включают:

  • model: сама модель PyTorch.
  • args: входной тензор или кортеж тензоров.
  • f: путь к выходному файлу.
  • export_params: экспорт всех весов модели вместе с графом.
  • opset_version: версия ONNX (рекомендуется использовать актуальную, например 17).
  • do_constant_folding: включение оптимизации констант.
  • input_names и output_names: читаемые имена входов и выходов.

Пример:

torch.onnx.export(
    model,
    dummy_input,
    "resnet18.onnx",
    export_params=True,
    opset_version=17,
    do_constant_folding=True,
    input_names=['input'],
    output_names=['output']
)

Проверка экспортированного графа

Для контроля корректности экспорта можно использовать модуль onnx для загрузки и проверки модели:

import onnx

onnx_model = onnx.load("resnet18.onnx")
onnx.checker.check_model(onnx_model)

Эта проверка выявляет ошибки совместимости и несоответствия версий операторов.

Оптимизация модели для веб

ONNX Runtime Web выполняет модели в браузере, и важно, чтобы граф был оптимизирован:

  • Удаление ненужных слоёв: слои, используемые только в обучении (например, Dropout), должны быть отключены.
  • Фиксация размеров входа: динамические размеры могут не поддерживаться или снижать производительность.
  • Выбор подходящей версии opset: современные браузеры и ONNX Runtime Web лучше работают с последними версиями.

Сохранение в формате совместимом с ONNX Runtime Web

Файл .onnx должен храниться в статическом виде и быть доступным для загрузки через HTTP. Рекомендуется проверить размер и структуру модели с помощью Netron или аналогичных инструментов для визуализации графа.

Особенности экспорта специфических слоёв

  • Conv, Linear: экспортируются без изменений.
  • BatchNorm: при режиме eval() объединяет параметры в константы.
  • RNN, LSTM, GRU: требуют фиксации состояния и последовательностей.
  • Custom Layers: необходимо реализовать через torch.autograd.Function с поддержкой ONNX или заменить стандартными эквивалентами.

Общая стратегия

  1. Подготовить модель в режиме eval().
  2. Сформировать пример входного тензора.
  3. Вызвать torch.onnx.export с настройкой имен, версии opset и оптимизаций.
  4. Проверить корректность экспортированного графа.
  5. При необходимости оптимизировать модель перед загрузкой в браузер.

Экспорт модели PyTorch в ONNX обеспечивает совместимость с ONNX Runtime Web, создавая основу для эффективного выполнения нейронных сетей на клиентской стороне.