Классификация изображений: ResNet, EfficientNet, MobileNet

ONNX Runtime Web (ORT Web) предоставляет возможность выполнять модели машинного обучения непосредственно в браузере с использованием JavaScript. Библиотека поддерживает форматы моделей ONNX, обеспечивает оптимизацию выполнения и позволяет работать с различными устройствами, включая CPU и WebGL для GPU-ускорения. ORT Web ориентирован на клиентские приложения, где важны производительность и переносимость моделей без серверной инфраструктуры.

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

import * as ort from 'onnxruntime-web';

const session = await ort.InferenceSession.create('model.onnx', {
    executionProviders: ['wasm', 'webgl'], // Выбор движков исполнения
    graphOptimizationLevel: 'all'          // Оптимизация графа
});

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

  • executionProviders — определяет, на каком движке будет выполняться модель (wasm для CPU, webgl для GPU через WebGL).
  • graphOptimizationLevel — оптимизация графа модели для ускорения инференса.

Подготовка данных для классификации изображений

Для классификации изображений требуется преобразовать входные данные в формат, совместимый с ONNX. Обычно модели ResNet, EfficientNet и MobileNet принимают изображения размером 224×224 или 224×224×3 и нормализованные пиксели.

Пример подготовки изображения из HTML <canvas>:

function preprocessImage(canvas) {
    const ctx = canvas.getContext('2d');
    const imageData = ctx.getImageData(0, 0, canvas.width, canvas.height);
    const { data, width, height } = imageData;

    const float32Data = new Float32Array(width * height * 3);
    for (let i = 0; i < width * height; i++) {
        float32Data[i * 3 + 0] = (data[i * 4 + 0] / 255.0 - 0.485) / 0.229; // R
        float32Data[i * 3 + 1] = (data[i * 4 + 1] / 255.0 - 0.456) / 0.224; // G
        float32Data[i * 3 + 2] = (data[i * 4 + 2] / 255.0 - 0.406) / 0.225; // B
    }

    return new ort.Tensor('float32', float32Data, [1, 3, height, width]);
}

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

  • Изображение преобразуется в тензор с формой [batch_size, channels, height, width].
  • Используется стандартная нормализация для предобученных моделей ImageNet.

Инициализация и запуск инференса

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

const inputTensor = preprocessImage(canvas);
const feeds = { input: inputTensor };

const results = await session.run(feeds);
const outputTensor = results.output;

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

  • Ключ словаря feeds должен соответствовать имени входного узла модели.
  • Результат results представляет собой словарь тензоров с именами выходов модели.

Постобработка результатов

Модели классификации возвращают вероятности классов. Для определения наиболее вероятного класса применяется операция argmax:

function getTopPrediction(outputTensor) {
    const data = outputTensor.data;
    let maxIndex = 0;
    let maxValue = data[0];

    for (let i = 1; i < data.length; i++) {
        if (data[i] > maxValue) {
            maxValue = data[i];
            maxIndex = i;
        }
    }
    return { classIndex: maxIndex, probability: maxValue };
}

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

  • Для моделей с softmax на выходе значение соответствует вероятности класса.
  • Для точного отображения классов требуется словарь с метками ImageNet.

Особенности ResNet, EfficientNet и MobileNet

ResNet

  • Глубокие остаточные сети, устойчивы к деградации при увеличении числа слоев.
  • Обычно используются конфигурации ResNet-18, ResNet-50 и ResNet-101.
  • Предпочтительно применять для задач с высокой точностью при умеренной производительности.

EfficientNet

  • Архитектура с масштабированием ширины, глубины и разрешения.
  • Обеспечивает высокую точность при меньшем числе параметров.
  • Подходит для браузерного инференса на ограниченных ресурсах.

MobileNet

  • Оптимизированные для мобильных и веб-приложений сети.
  • Использует глубинные свертки (depthwise separable convolution) для снижения вычислительной нагрузки.
  • Хороший выбор для быстрого инференса на GPU через WebGL.

Оптимизация работы ONNX Runtime Web

  1. Выбор движка исполнения Использование webgl позволяет ускорить инференс на GPU, особенно для больших моделей. В случаях ограниченного GPU лучше использовать wasm.

  2. Минимизация передачи данных Избегать лишнего копирования изображений и данных между DOM и памятью JavaScript.

  3. Асинхронный инференс Выполнять инференс асинхронно, чтобы не блокировать основной поток браузера.

  4. Оптимизация модели Перед использованием модели в вебе рекомендуется применять оптимизации ONNX (onnxruntime-tools) для уменьшения размера и ускорения графа.

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

const canvas = document.getElementById('imageCanvas');
const session = await ort.InferenceSession.create('mobilenet.onnx', {
    executionProviders: ['webgl']
});

const inputTensor = preprocessImage(canvas);
const results = await session.run({ input: inputTensor });
const prediction = getTopPrediction(results.output);

console.log(`Предсказанный класс: ${prediction.classIndex}, вероятность: ${prediction.probability}`);

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