Стриминг загрузки модели с прогрессом

ONNX Runtime Web (ORT Web) предоставляет возможности для запуска моделей машинного обучения прямо в браузере с использованием WebAssembly (WASM) или WebGL. Одной из важных задач при работе с крупными моделями является организация эффективной загрузки с отображением прогресса, чтобы пользовательский интерфейс оставался отзывчивым и информативным.

Асинхронная загрузка модели

В ORT Web модели загружаются через InferenceSession.create() или InferenceSession.loadModel(). Для больших моделей это может занимать значительное время, особенно при использовании сети с ограниченной пропускной способностью. Основной подход — разделить процесс загрузки на части и отслеживать прогресс.

Пример базовой асинхронной загрузки:

import * as ort from 'onnxruntime-web';

async function loadModel(url) {
  const session = await ort.InferenceSession.create(url);
  return session;
}

Проблема такого подхода в том, что нет встроенного механизма уведомления о прогрессе. Для решения используют Streaming API браузера и fetch с чтением тела ответа по кускам.

Использование Fetch Streaming API

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

async function loadModelWithProgress(url, onProgress) {
  const response = await fetch(url);
  const contentLength = response.headers.get('Content-Length');

  if (!contentLength) {
    throw new Error('Не удалось определить размер файла для прогресса');
  }

  const total = parseInt(contentLength, 10);
  let loaded = 0;

  const reader = response.body.getReader();
  const chunks = [];

  while (true) {
    const { done, value } = await reader.read();
    if (done) break;
    chunks.push(value);
    loaded += value.length;
    onProgress(loaded / total);
  }

  const modelArray = new Uint8Array(loaded);
  let position = 0;
  for (const chunk of chunks) {
    modelArray.set(chunk, position);
    position += chunk.length;
  }

  const session = await ort.InferenceSession.create(modelArray);
  return session;
}

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

  • response.body.getReader() позволяет читать поток данных по частям.
  • value.length измеряет количество загруженных байт.
  • onProgress — callback-функция для обновления интерфейса загрузки, принимает дробное число от 0 до 1.
  • Сборка массива Uint8Array из всех чанков обеспечивает совместимость с InferenceSession.create.

Интеграция с пользовательским интерфейсом

Прогресс загрузки можно визуализировать через прогресс-бары или текстовые индикаторы. Например:

function updateProgressBar(fraction) {
  const progressBar = document.getElementById('progress-bar');
  progressBar.value = fraction * 100;
}

И использование вместе с загрузкой:

const session = await loadModelWithProgress('model.onnx', updateProgressBar);
console.log('Модель загружена', session);

Оптимизация загрузки больших моделей

  1. Сжатие модели: gzip или brotli может существенно уменьшить размер передаваемых данных. В браузере можно использовать DecompressionStream для распаковки на лету.
  2. Использование CDN: размещение модели на распределённом сервере ускоряет потоковую передачу.
  3. Пакетная загрузка: разделение модели на несколько частей и последующая сборка в Uint8Array позволяет показывать прогресс для каждой части отдельно.
  4. Параллельные запросы: для больших моделей можно загружать несколько чанков параллельно и собирать их, однако требует контроля порядка байт в итоговом массиве.

Работа с WASM и WebGL бэкендами

ORT Web поддерживает разные бэкенды исполнения:

  • WASM — универсальный, работает во всех современных браузерах.
  • WebGL / WebGPU — ускорение через GPU, особенно полезно для больших сетей.

Различие в загрузке модели минимальное, однако при использовании GPU-ускорения важно ожидать инициализацию контекста WebGL, чтобы модель сразу могла быть использована без задержек:

await ort.env.warmUp();
const session = await loadModelWithProgress('model.onnx', updateProgressBar);

Отладка и логирование

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

function logProgress(fraction) {
  console.log(`Загружено: ${(fraction * 100).toFixed(2)}%`);
}

const session = await loadModelWithProgress('model.onnx', logProgress);

Такой подход помогает выявлять узкие места в сети или сервере и оптимизировать UX при загрузке моделей свыше 100 МБ.

Практические рекомендации

  • Всегда проверять наличие заголовка Content-Length, иначе прогресс вычислить невозможно.
  • Использовать Uint8Array вместо массивов или строк для бинарных данных.
  • Реализовывать тайм-ауты на fetch-запросы для предотвращения зависания при медленном соединении.
  • Для крупных моделей комбинировать потоковую загрузку с индикаторами прогресса на основе событий fetch и Reader.

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