Отладка NaN и Inf в выходных тензорах

Работа с ONNX Runtime Web в JavaScript нередко сопровождается ситуациями, когда модель возвращает значения NaN (Not a Number) или Inf (Infinity) в выходных тензорах. Такие значения могут свидетельствовать о нестабильности модели, ошибках данных или некорректной настройке вычислений. Понимание их происхождения и эффективная диагностика критически важны для обеспечения корректного функционирования приложений.


Причины появления NaN и Inf

1. Некорректные входные данные

  • Пустые значения, undefined или null в тензорах.
  • Слишком большие или слишком маленькие значения, выходящие за диапазон допустимых значений для чисел с плавающей точкой.
  • Некорректная нормализация данных (например, деление на ноль).

2. Ошибки модели или её конфигурации

  • Некорректные веса модели, полученные при экспорте в ONNX.
  • Операции, приводящие к переполнению или делению на ноль (exp, log, div).
  • Проблемы с типами данных тензоров (float32 vs float64).

3. Особенности выполнения в браузере

  • Ограничения производительности и точности вычислений на WebGL или WebAssembly.
  • Потеря точности при конвертации данных между форматами JavaScript и ONNX.

Методы обнаружения NaN и Inf

1. Использование функций проверки JavaScript

Для проверки отдельных элементов тензора применяются стандартные методы:

function hasNaN(tensorData) {
    return tensorData.some(x => Number.isNaN(x));
}

function hasInf(tensorData) {
    return tensorData.some(x => !Number.isFinite(x));
}

При больших тензорах стоит использовать TypedArray методы для оптимизации:

const data = outputTensor.data;
const hasNaN = data.some(Number.isNaN);
const hasInf = data.some(x => !Number.isFinite(x));

2. Отладка на уровне сессии ONNX Runtime

ONNX Runtime Web предоставляет возможность отладки сессий и слоёв модели:

const session = await ort.InferenceSession.create('model.onnx', { executionProviders: ['wasm'] });
const feeds = { input: new ort.Tensor('float32', inputData, inputShape) };
const results = await session.run(feeds);
console.log(results.outputTensor.data.slice(0, 10));

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


Методы локализации проблем

1. Разделение входов и слоёв

  • Проверять корректность каждого входного тензора.
  • Прогонять модель по слоям (если модель небольшая или экспортирована с поддержкой nodejs для детальной трассировки).

2. Использование промежуточных тензоров

  • Включение выводов промежуточных слоёв при экспорте модели.
  • Сравнение значений между слоями, выявление первого появления NaN или Inf.

3. Нормализация и масштабирование данных

  • Передача входных данных через функции нормализации:
function normalize(input) {
    const max = Math.max(...input);
    const min = Math.min(...input);
    return input.map(x => (x - min) / (max - min));
}
  • Ограничение значений до диапазона [minThreshold, maxThreshold] для предотвращения переполнения.

Превентивные меры

  • Использование типа float32 для всех тензоров.
  • Проверка модели перед экспортом в ONNX с использованием тестовых данных.
  • Избегание операций, потенциально приводящих к Inf, без предварительной защиты (например, log(0) или деление на ноль).

Практический пример: фильтрация NaN и Inf

function sanitizeTensorData(data) {
    return data.map(x => {
        if (Number.isNaN(x)) return 0;
        if (!Number.isFinite(x)) return x > 0 ? Number.MAX_VALUE : -Number.MAX_VALUE;
        return x;
    });
}

const sanitizedData = sanitizeTensorData(outputTensor.data);
const sanitizedTensor = new ort.Tensor('float32', sanitizedData, outputTensor.dims);

Такой подход предотвращает распространение некорректных значений на последующие вычисления.


Логирование и визуализация

  • Использование console.table или графических библиотек для визуального контроля распределения значений.
  • Вывод статистики: максимальные, минимальные значения, количество NaN/Inf.
const stats = {
    max: Math.max(...outputTensor.data),
    min: Math.min(...outputTensor.data),
    NaN_count: outputTensor.data.filter(Number.isNaN).length,
    Inf_count: outputTensor.data.filter(x => !Number.isFinite(x)).length
};
console.table(stats);

Взаимодействие с WebGL и WebAssembly

  • При использовании WebGL иногда возникают неявные переполнения или округления.
  • Настройка провайдера выполнения (wasm или webgl) позволяет сравнивать результаты и выявлять платформенные особенности.
const sessionWASM = await ort.InferenceSession.create('model.onnx', { executionProviders: ['wasm'] });
const sessionWebGL = await ort.InferenceSession.create('model.onnx', { executionProviders: ['webgl'] });

Сравнение выходов помогает локализовать источник некорректных значений.


Рекомендации по оптимизации

  • Разделение больших тензоров на блоки при обработке для уменьшения ошибок переполнения.
  • Использование промежуточного контроля значений после каждой крупной операции.
  • Поддержка автоматической очистки NaN и Inf до передачи данных на последующие слои или в интерфейс.