Экспорт из scikit-learn через sklearn-onnx

Для работы с ONNX Runtime Web в JavaScript необходимо подготовить окружение, включающее Node.js или браузерный контекст, а также пакеты для взаимодействия с ONNX-моделями. Для работы в Node.js требуется установка onnxruntime-web через npm:

npm install onnxruntime-web

Для экспорта моделей из scikit-learn используется пакет sklearn-onnx:

pip install skl2onnx onnx

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


Экспорт модели из scikit-learn в ONNX

  1. Создание и обучение модели:
from sklearn.datasets import load_iris
from sklearn.ensemble import RandomForestClassifier
from sklearn.model_selection import train_test_split

iris = load_iris()
X_train, X_test, y_train, y_test = train_test_split(iris.data, iris.target, test_size=0.2)

model = RandomForestClassifier(n_estimators=100)
model.fit(X_train, y_train)
  1. Экспорт модели в ONNX:

Для корректного экспорта необходимо задать типы входных данных через FloatTensorType.

from skl2onnx import convert_sklearn
from skl2onnx.common.data_types import FloatTensorType

initial_type = [('float_input', FloatTensorType([None, X_train.shape[1]]))]
onnx_model = convert_sklearn(model, initial_types=initial_type)

with open("rf_model.onnx", "wb") as f:
    f.write(onnx_model.SerializeToString())

Ключевой момент: initial_type задаёт структуру входных данных и обязателен для корректного использования модели в ONNX Runtime Web. None обозначает динамическое количество примеров, что удобно для обработки пакетов данных в JavaScript.


Импорт и использование модели в JavaScript

  1. Загрузка модели:
import * as ort from 'onnxruntime-web';

async function loadModel() {
    const session = await ort.InferenceSession.create('rf_model.onnx');
    return session;
}
  1. Подготовка входных данных:

Входные данные передаются в формате TypedArray или Tensor от ONNX Runtime:

const input = new Float32Array([5.1, 3.5, 1.4, 0.2]); // Пример одного наблюдения
const tensor = new ort.Tensor('float32', input, [1, 4]); // [batch, features]

Ключевой момент: размерность тензора должна точно соответствовать initial_type, указанному при экспорте модели. Несоответствие вызовет ошибку при инференсе.

  1. Запуск инференса:
async function predict(session, tensor) {
    const feeds = { float_input: tensor };
    const results = await session.run(feeds);
    console.log(results);
}

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


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

ONNX Runtime Web поддерживает батчинг, что позволяет обрабатывать несколько наблюдений одновременно:

const batchInput = new Float32Array([
    5.1, 3.5, 1.4, 0.2,
    6.2, 3.4, 5.4, 2.3
]); // 2 наблюдения, 4 признака
const batchTensor = new ort.Tensor('float32', batchInput, [2, 4]);

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

  • Использовать WebAssembly (wasm) backend или WebGPU при наличии GPU.
  • Загружать модель один раз и повторно использовать с разными входами.
  • Приводить все входные данные к нужному типу заранее (Float32Array) для ускорения инференса.

Совместимость и ограничения

  • Не все алгоритмы scikit-learn имеют полную поддержку в ONNX. Например, сложные пайплайны с кастомными трансформерами могут потребовать ручной сериализации или переписывания на совместимые элементы.
  • Категориальные данные должны быть закодированы численно (LabelEncoder, OneHotEncoder) перед экспортом.
  • ONNX Runtime Web работает как в Node.js, так и в браузере, но в браузере доступны только WebAssembly и WebGPU бэкенды, что может влиять на скорость инференса для больших моделей.

Интеграция в фронтенд-приложения

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

Пример асинхронного предсказания в браузере:

document.getElementById('predict-btn').addEventListener('click', async () => {
    const session = await loadModel();
    const inputValues = new Float32Array([5.1, 3.5, 1.4, 0.2]);
    const tensor = new ort.Tensor('float32', inputValues, [1, 4]);
    const results = await predict(session, tensor);
    console.log('Prediction:', results);
});

Ключевой момент: асинхронная загрузка позволяет не блокировать UI и безопасно работать с большими моделями.