Несоответствие типов и форм тензоров

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


Типы тензоров

Каждый вход модели в ONNX имеет строго определённый тип данных. Типы тензоров в ONNX включают:

  • float32 — наиболее часто используемый тип для числовых данных с плавающей точкой.
  • int32, int64 — целочисленные типы для категориальных или индексных данных.
  • bool — булевы значения, например для масок.
  • string — строки, применяются редко и требуют отдельной обработки.

Несоответствие типа возникает, если передать тензор другого типа. Например, если модель ожидает float32, а передан int32, ORT Web выбросит ошибку на этапе session.run() с сообщением о несоответствии типов.

const inputTensor = new ort.Tensor('int32', new Int32Array([1, 2, 3]), [3]);

Если модель ожидает float32, необходимо привести данные:

const inputTensor = new ort.Tensor('float32', new Float32Array([1.0, 2.0, 3.0]), [3]);

Выделение ключевых моментов:

  • Типы тензоров строго проверяются.
  • Приведение типов должно выполняться вручную, ORT Web не конвертирует данные автоматически.
  • Ошибки типов часто встречаются при работе с массивами JavaScript, так как Array по умолчанию не имеет строгого типа.

Формы (shapes) тензоров

Форма тензора — это массив, определяющий размеры по каждой оси. Например, [1, 3, 224, 224] — это тензор для изображения с одной картинкой, 3 каналами, размером 224×224.

Основные правила:

  • Количество элементов в массиве данных должно соответствовать произведению всех размеров формы.
  • Несоответствие формы вызывает ошибку при вызове session.run().

Пример ошибки формы:

const inputTensor = new ort.Tensor('float32', new Float32Array(150528), [1, 3, 224, 223]);
// Ошибка: количество элементов (150528) не соответствует форме [1, 3, 224, 223] (1*3*224*223 = 149,856)

Правильный вариант:

const inputTensor = new ort.Tensor('float32', new Float32Array(150528), [1, 3, 224, 224]);

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

  • Всегда проверять произведение всех размеров формы против длины массива данных.
  • Использование console.log(inputTensor.dims) помогает выявлять несоответствия.
  • Динамические формы (-1 в ONNX) требуют дополнительной логики для вычисления конкретных размеров.

Проверка и трансформация входных данных

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

function createTensor(data, expectedType, expectedShape) {
    let typedArray;
    switch (expectedType) {
        case 'float32': typedArray = new Float32Array(data); break;
        case 'int32': typedArray = new Int32Array(data); break;
        case 'int64': typedArray = BigInt64Array.from(data.map(BigInt)); break;
        default: throw new Error(`Unsupported tensor type: ${expectedType}`);
    }

    const expectedLength = expectedShape.reduce((a, b) => a * b, 1);
    if (typedArray.length !== expectedLength) {
        throw new Error(`Tensor length ${typedArray.length} does not match expected shape ${expectedShape}`);
    }

    return new ort.Tensor(expectedType, typedArray, expectedShape);
}

Преимущества такой проверки:

  • Исключение ошибок на этапе подготовки данных.
  • Автоматическое приведение типов и проверка формы.
  • Удобство работы с динамическими формами.

Особенности работы с батчами

Многие модели требуют входного тензора с размерностью [batch, channels, height, width]. Частой ошибкой является передача тензора без батча, например [3, 224, 224] вместо [1, 3, 224, 224]. Для ORT Web это критично, так как библиотека строго соблюдает спецификацию ONNX.

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

const batchTensor = new ort.Tensor('float32', new Float32Array(1 * 3 * 224 * 224), [1, 3, 224, 224]);

Вывод: даже если данные визуально корректны, отсутствие первой оси батча может привести к ошибке выполнения.


Ошибки при несоответствии типов и форм

Типичные сообщения ORT Web:

  • TypeError: Expected input type float32 but received int32
  • ShapeError: Tensor length does not match shape
  • RuntimeError: Invalid input shape

Причины ошибок:

  1. Использование стандартного JavaScript Array вместо TypedArray.
  2. Неверное умножение размеров формы и длины данных.
  3. Пропущенная первая размерность батча.
  4. Несоответствие типов при передаче индексов или категориальных данных.

Правильная подготовка тензоров минимизирует такие ошибки и позволяет ORT Web эффективно выполнять вычисления без аварийных остановок.


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