Для работы с 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, что
включает такие алгоритмы, как линейная регрессия, логистическая
регрессия, деревья решений, случайные леса и градиентный бустинг.
Поддержка некоторых нестандартных препроцессоров может потребовать
ручной настройки преобразований.
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)
Для корректного экспорта необходимо задать типы входных данных через
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.
import * as ort from 'onnxruntime-web';
async function loadModel() {
const session = await ort.InferenceSession.create('rf_model.onnx');
return session;
}
Входные данные передаются в формате 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, указанному при экспорте
модели. Несоответствие вызовет ошибку при инференсе.
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]);
Оптимизация производительности:
wasm) backend или WebGPU при
наличии GPU.Float32Array) для ускорения инференса.scikit-learn имеют полную поддержку в
ONNX. Например, сложные пайплайны с кастомными трансформерами могут
потребовать ручной сериализации или переписывания на совместимые
элементы.LabelEncoder, OneHotEncoder) перед
экспортом.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 и безопасно работать с большими моделями.