Экспорт модели из TensorFlow и Keras

ONNX Runtime Web (ORT Web) предоставляет возможность запуска моделей машинного обучения в браузере на JavaScript. Для использования ORT Web сначала необходимо конвертировать модели, созданные в TensorFlow или Keras, в формат ONNX. Процесс экспорта требует соблюдения определённых правил совместимости и внимательного подхода к подготовке модели.

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

Перед экспортом необходимо убедиться, что модель завершена и оптимизирована:

  • Фиксированная форма входных данных: ORT Web требует статических размеров тензоров, поэтому динамические размеры должны быть зафиксированы.
  • Удаление слоёв, несовместимых с ONNX: Некоторые специфические слои Keras или TensorFlow могут не поддерживаться. В таких случаях следует заменить их на эквиваленты, совместимые с ONNX, или реализовать через кастомные операции.
  • Компиляция модели: Для Keras рекомендуется сохранить модель с compile=False, если нет необходимости в сохранении метрик и оптимизаторов, так как они не используются ONNX.

Пример подготовки модели в Keras:

import tensorflow as tf
from tensorflow.keras.models import Sequential
from tensorflow.keras.layers import Dense

model = Sequential([
    Dense(64, activation='relu', input_shape=(100,)),
    Dense(10, activation='softmax')
])
# Компиляция не обязательна для экспорта
model.compile(optimizer='adam', loss='categorical_crossentropy')

Конвертация в ONNX

Для конвертации TensorFlow/Keras моделей в ONNX существует несколько инструментов, в том числе tf2onnx и keras2onnx. Основные шаги:

  1. Установка инструментов:
pip install tf2onnx onnx
  1. Конвертация модели:
python -m tf2onnx.convert --saved-model ./saved_model \
                           --output model.onnx \
                           --opset 13
  • --saved-model указывает путь к сохранённой модели TensorFlow.
  • --opset определяет версию ONNX, совместимую с ORT Web (рекомендуется 13–15).

Для моделей Keras:

import keras2onnx
import onnx

onnx_model = keras2onnx.convert_keras(model, model.name)
onnx.save_model(onnx_model, "model.onnx")

Проверка совместимости

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

import onnx

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

Проверка выявляет ошибки совместимости, такие как несовпадение типов тензоров или некорректные размеры. Исправление проблем на этом этапе предотвращает ошибки при загрузке модели в ORT Web.

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

Для работы в браузере следует минимизировать размер модели и количество операций:

  • Удаление ненужных слоёв: Например, слоёв, используемых только во время обучения (Dropout).
  • Фиксация весов и параметров: Все параметры должны быть константами, чтобы ORT Web не требовал дополнительных вычислений.
  • Сжатие модели: Использование инструментов типа onnxoptimizer или квантование до INT8, что снижает размер файла и ускоряет выполнение в браузере.

Загрузка и использование ONNX модели в ORT Web

ORT Web предоставляет объект InferenceSession для работы с моделью:

import * as ort from "onnxruntime-web";

async function runModel() {
    const session = await ort.InferenceSession.create('model.onnx');
    const inputTensor = new ort.Tensor('float32', inputData, [1, 100]);
    const feeds = { input_name: inputTensor };
    const results = await session.run(feeds);
    console.log(results.output_name.data);
}

Ключевые моменты:

  • input_name и output_name должны соответствовать именам тензоров в ONNX-модели.
  • Типы данных тензоров должны совпадать с типами модели (float32, int32 и т.д.).
  • Размерность входного тензора должна строго соответствовать статической форме, заданной при конвертации.

Советы по дебагу

  • Использовать onnxruntime-web с логированием для отслеживания ошибок загрузки.
  • Проверять веса и форму тензоров через Netron или аналогичный инструмент визуализации ONNX.
  • Проверять поддержку используемых операций в ORT Web, так как некоторые сложные операции могут быть реализованы частично или требовать fallback на CPU.

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