Загрузка графовых моделей через tf.loadGraphModel

tf.loadGraphModel — это ключевой метод библиотеки TensorFlow.js, предназначенный для загрузки заранее обученных моделей в формате графа вычислений (Graph Model). В отличие от моделей последовательного типа (Sequential), графовые модели представляют собой более сложные архитектуры с возможностью наличия ветвлений, объединений и произвольных связей между слоями.


Формат графовой модели

Графовые модели TensorFlow.js обычно хранятся в формате JSON, сопровождаемом бинарными файлами весов. Структура выглядит следующим образом:

model.json        // Основной файл модели с описанием графа и метаинформацией
group1-shard1.bin // Бинарные файлы весов
group1-shard2.bin
...

Файл model.json содержит:

  • Описание слоев и операций;
  • Информацию о входных и выходных тензорах;
  • Ссылки на бинарные файлы с весами (weightsManifest).

Это позволяет браузеру или Node.js динамически подгружать веса и конфигурацию без необходимости конвертации в другой формат.


Синтаксис метода

const model = await tf.loadGraphModel(modelUrl, options);

Параметры:

  • modelUrl — строка с URL или локальным путем к файлу model.json. Поддерживаются:

    • HTTP(S) URL;
    • Локальные файлы в Node.js через file://;
    • Файлы в IndexedDB (indexeddb://my-model).
  • options (необязательный объект):

    • fromTFHub (boolean) — указывает, что модель загружается с TensorFlow Hub; автоматически применяет необходимые преобразования URL.
    • requestInit — объект для настройки HTTP-запроса (например, заголовки, метод, credentials).

Пример загрузки из URL:

const model = await tf.loadGraphModel('https://example.com/model/model.json');

Пример загрузки из IndexedDB:

const model = await tf.loadGraphModel('indexeddb://my-model');

Работа с входами и выходами модели

Графовые модели принимают данные в виде объектов tf.Tensor или массивов тензоров. Для выполнения предсказания используется метод model.execute или model.predict.

model.execute позволяет управлять конкретными входами и выходами:

const inputTensor = tf.tensor2d([[1, 2, 3, 4]], [1, 4]);
const outputTensor = model.execute({ 'input_1': inputTensor }, 'output_node');
outputTensor.print();
  • Первый аргумент — объект, где ключи соответствуют именам входных узлов графа.
  • Второй аргумент — имя выходного узла или массив имен узлов, результаты которых необходимо получить.

model.predict работает, если модель имеет единственный вход и один выход:

const inputTensor = tf.tensor2d([[1, 2, 3, 4]], [1, 4]);
const outputTensor = model.predict(inputTensor);
outputTensor.print();

Асинхронная природа загрузки

tf.loadGraphModel всегда возвращает промис, что требует использования await или then. Это связано с загрузкой бинарных весов через сеть или локальное хранилище. В случае больших моделей важно использовать индикаторы загрузки и освобождение ресурсов после использования:

const model = await tf.loadGraphModel('https://example.com/model/model.json');

// Освобождение ресурсов
model.dispose();

Оптимизация и кэширование

TensorFlow.js автоматически кэширует загруженные модели в памяти. Для повторного использования можно:

  • Загружать модель один раз и сохранять ссылку в переменной.
  • Сохранять модель в IndexedDB для быстрого доступа при следующих сессиях:
await model.save('indexeddb://my-cached-model');

Загрузка из IndexedDB ускоряет процесс, так как не требуется повторная загрузка через сеть.


Обработка ошибок при загрузке

Основные причины ошибок:

  1. Неверный путь к model.json — возникает 404 или ошибка сети.
  2. Несовместимый формат — загружается не TensorFlow.js Graph Model.
  3. Проблемы с весами — отсутствует или поврежден .bin файл.

Пример обработки ошибок:

try {
    const model = await tf.loadGraphModel('https://example.com/model/model.json');
    console.log('Модель успешно загружена');
} catch (err) {
    console.error('Ошибка загрузки модели:', err);
}

Преобразование моделей TensorFlow и Keras

Модели TensorFlow SavedModel или Keras можно конвертировать в формат TensorFlow.js с помощью утилиты tensorflowjs_converter. Пример команды:

tensorflowjs_converter \
  --input_format=tf_saved_model \
  --output_format=tfjs_graph_model \
  /path/to/saved_model \
  /path/to/tfjs_model

После этого можно загружать модель в браузере или Node.js через tf.loadGraphModel.


Важные рекомендации по использованию

  • Использовать явные имена входных и выходных узлов для предотвращения ошибок при сложных графах.
  • Очистка памяти через dispose или tf.tidy после работы с тензорами.
  • Для больших моделей — предварительно загрузка весов и прогрев модели, чтобы снизить задержку первых предсказаний.
  • Проверка совместимости с бэкендом (tf.setBackend('webgl') для ускорения вычислений в браузере).

Пример интеграции в веб-приложение

<script src="https://cdn.jsdelivr.net/npm/@tensorflow/tfjs"></script>
<script>
async function runModel() {
    const model = await tf.loadGraphModel('https://example.com/model/model.json');

    const input = tf.tensor2d([[0.5, -1.2, 3.3]], [1, 3]);
    const output = model.execute({ 'input_node': input }, 'output_node');
    output.print();

    model.dispose();
}
runModel();
</script>

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


Механизм tf.loadGraphModel обеспечивает гибкость и масштабируемость при работе с сложными нейронными сетями, позволяя интегрировать модели TensorFlow в веб и серверные приложения на JavaScript с высокой производительностью и контролем над ресурсами.