Батчирование: формирование батча из нескольких входов

Батчирование является одной из ключевых техник оптимизации при работе с нейронными сетями. В контексте ONNX Runtime Web (ORT Web) оно позволяет объединять несколько входных данных в один пакет, сокращая количество вызовов модели и повышая общую производительность вычислений.

Основные принципы батчирования

Батч представляет собой многомерный тензор, где первая размерность отвечает за количество элементов в пакете. Для модели, принимающей на вход тензор формы [N, C, H, W], батчирование позволяет подать сразу N изображений, вместо последовательной обработки каждого из них.

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

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

Подготовка входных данных

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

import * as ort from 'onnxruntime-web';

// Функция для нормализации и конвертации изображения в Float32Array
function preprocessImage(image) {
    // image: HTMLImageElement или ImageBitmap
    const canvas = document.createElement('canvas');
    canvas.width = image.width;
    canvas.height = image.height;
    const ctx = canvas.getContext('2d');
    ctx.drawImage(image, 0, 0);
    const imageData = ctx.getImageData(0, 0, image.width, image.height);
    const data = imageData.data;
    const floatData = new Float32Array(image.width * image.height * 3);
    for (let i = 0; i < image.width * image.height; i++) {
        floatData[i * 3 + 0] = data[i * 4 + 0] / 255.0;
        floatData[i * 3 + 1] = data[i * 4 + 1] / 255.0;
        floatData[i * 3 + 2] = data[i * 4 + 2] / 255.0;
    }
    return floatData;
}

После подготовки всех отдельных изображений нужно объединить их в единый батч:

function createBatch(images, width, height) {
    const batchSize = images.length;
    const batchData = new Float32Array(batchSize * 3 * height * width);
    for (let i = 0; i < batchSize; i++) {
        const imageData = preprocessImage(images[i]);
        batchData.set(imageData, i * 3 * height * width);
    }
    return new ort.Tensor('float32', batchData, [batchSize, 3, height, width]);
}

Использование батча с ONNX Runtime Web

После формирования батча создается сессия и выполняется инференс:

async function runBatchInference(session, batchTensor) {
    const feeds = { input: batchTensor }; // 'input' — имя входного узла модели
    const results = await session.run(feeds);
    return results.output; // 'output' — имя выходного узла модели
}

Особенности работы с ORT Web:

  • Для больших батчей может увеличиваться использование памяти браузера.
  • ORT Web поддерживает асинхронные вызовы, что важно при работе с батчами для предотвращения блокировки UI.
  • Размер батча можно варьировать динамически в зависимости от производительности устройства.

Оптимизация батчирования

  • Минимизация лишних копирований данных. Создание Float32Array напрямую в нужной размерности снижает накладные расходы.
  • Предварительное выделение памяти под максимальный размер батча ускоряет повторные вызовы.
  • Выравнивание формы данных. Если модель ожидает конкретную форму [N, C, H, W], необходимо обеспечить одинаковый размер каждого элемента батча.
  • Динамический батчинг. В веб-приложениях можно накапливать несколько запросов пользователей и объединять их в один батч, повышая пропускную способность модели.

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

let pendingImages = [];
const MAX_BATCH_SIZE = 8;

function enqueueImage(image, session) {
    pendingImages.push(image);
    if (pendingImages.length >= MAX_BATCH_SIZE) {
        const batchTensor = createBatch(pendingImages, 224, 224);
        runBatchInference(session, batchTensor).then(output => {
            console.log('Batch output:', output.data);
        });
        pendingImages = [];
    }
}

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

Совместимость типов данных

Модели ONNX строго типизированы. При формировании батча следует использовать тот же тип, что и у входного узла модели (float32, int64 и т. д.). Несоответствие типов приведет к ошибкам выполнения:

// Ошибочно:
const batchTensor = new ort.Tensor('int32', batchData, [batchSize, 3, 224, 224]);
// Правильно:
const batchTensor = new ort.Tensor('float32', batchData, [batchSize, 3, 224, 224]);

Важные рекомендации

  • Проверять соответствие размерностей батча и формы входного узла модели.
  • Измерять производительность при разных размерах батча для нахождения оптимального баланса между временем инференса и использованием памяти.
  • При использовании WebGL или WebAssembly backend учитывать ограничения GPU и браузера по объему доступной памяти.

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