Регрессия и табличные данные: деревья решений, линейные модели

ONNX Runtime Web (ORT Web) представляет собой высокопроизводительную среду выполнения моделей машинного обучения в браузере и в Node.js на базе формата ONNX. Этот подход позволяет интегрировать модели регрессии и классификации, обученные на Python или других платформах, в JavaScript-приложения без необходимости запускать серверную часть. Основной принцип работы заключается в загрузке модели ONNX и выполнении предсказаний на клиентской стороне.

Установка и подключение

Для использования ORT Web в браузере подключение осуществляется через npm-пакет:

import * as ort from 'onnxruntime-web';

или через тег <script>:

<script src="https://cdn.jsdelivr.net/npm/onnxruntime-web/dist/ort.min.js"></script>

Для Node.js рекомендуется использовать пакет onnxruntime-node для оптимальной производительности, особенно при работе с большими моделями регрессии.

Загрузка модели

Модель ONNX загружается с помощью метода InferenceSession.create, который поддерживает указание как локального файла, так и URL-адреса:

const session = await ort.InferenceSession.create('model.onnx');

Параметры конфигурации включают выбор оптимизатора и использование WebAssembly (WASM) или WebGL для ускорения вычислений:

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

Подготовка данных для регрессии

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

const inputData = new Float32Array([5.1, 3.5, 1.4, 0.2]);
const feeds = { input: new ort.Tensor('float32', inputData, [1, 4]) };
  • input — имя входного узла модели, которое задаётся при экспорте модели в ONNX.
  • [1, 4] — форма тензора (batch size = 1, 4 признака).

Для табличных данных с несколькими объектами создаётся двумерный тензор [n, m], где n — количество строк, а m — количество признаков.

Выполнение предсказаний

Для получения результатов используется метод session.run(feeds):

const results = await session.run(feeds);
const output = results.output.data;
console.log(output);
  • results.output.data возвращает одномерный массив чисел для регрессии.
  • В случае деревьев решений и ансамблей (Random Forest, Gradient Boosting) значения могут быть суммой предсказаний отдельных деревьев.

Особенности линейных моделей

Линейные модели регрессии обладают высокой скоростью работы и прозрачностью. Основные шаги подготовки данных включают нормализацию и кодирование категориальных признаков:

  • Нормализация: стандартное масштабирование признаков повышает точность предсказаний.
  • One-hot encoding: преобразование категориальных переменных в числовые тензоры.

Пример кода для подготовки данных с нормализацией:

const mean = [5.0, 3.4, 1.5, 0.2];
const std = [0.5, 0.3, 0.2, 0.1];
const normalized = inputData.map((val, idx) => (val - mean[idx]) / std[idx]);
const tensor = new ort.Tensor('float32', normalized, [1, 4]);

Особенности деревьев решений

Деревья решений и ансамбли требуют небольших модификаций в данных:

  • Отсутствие необходимости масштабирования числовых признаков.
  • Поддержка категориальных признаков через one-hot encoding или прямое кодирование.
  • Возможность предсказаний сразу для нескольких объектов (batch inference).

Пример предсказания для нескольких объектов:

const batchData = new Float32Array([
  5.1, 3.5, 1.4, 0.2,
  6.2, 3.4, 5.4, 2.3
]);
const feeds = { input: new ort.Tensor('float32', batchData, [2, 4]) };
const results = await session.run(feeds);
console.log(results.output.data);

Оптимизация производительности

При работе с большими табличными наборами данных и ансамблями деревьев решений важны следующие аспекты:

  • Использование WebGL для ускорения матричных вычислений.
  • Пакетная обработка данных: передача нескольких строк одновременно минимизирует накладные расходы на создание тензоров.
  • Минимизация копирования данных: применение Float32Array вместо обычных массивов JavaScript.

Управление памятью и потоками

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

session.dispose();

При необходимости многократных предсказаний рекомендуется повторно использовать объект session для экономии времени на инициализацию модели.

Интеграция с UI

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

  • Построение графиков предсказаний.
  • Реализация интерактивных таблиц с вводом новых данных.
  • Использование библиотек визуализации, таких как Chart.js или D3.js, для отображения трендов.

Работа с вероятностными предсказаниями

Модели деревьев решений, такие как Random Forest и Gradient Boosting, могут возвращать вероятности классов. В регрессии это может быть полезно для построения доверительных интервалов:

const results = await session.run(feeds);
const predictions = results.output.data;
const confidence = results.prob.data; // при наличии соответствующего узла в модели

Рекомендации по экспорту моделей в ONNX

  • Для линейной регрессии в Python использовать sklearn.linear_model.LinearRegression с экспортом через skl2onnx.
  • Для деревьев решений и ансамблей использовать sklearn.tree.DecisionTreeRegressor и sklearn.ensemble с последующим сохранением через skl2onnx.
  • Проверять имена узлов входа и выхода, чтобы корректно создавать тензоры в JavaScript.

Эта методика обеспечивает переносимость моделей между средами и упрощает интеграцию регрессионных моделей и деревьев решений с веб-приложениями.