Экспорт модели в ONNX-совместимый формат

Brain.js — это библиотека для работы с нейронными сетями в среде JavaScript, которая позволяет создавать и обучать модели для различных задач, таких как классификация, регрессия и распознавание паттернов. Одной из важных возможностей является экспорт обученной модели для использования вне среды Node.js или браузера, например, в системах, поддерживающих стандарт ONNX (Open Neural Network Exchange). ONNX обеспечивает совместимость моделей между различными фреймворками, такими как PyTorch, TensorFlow и другим ПО для машинного обучения.

Подготовка модели к экспорту

Перед экспортом необходимо убедиться, что модель обучена и готова к использованию. Brain.js предоставляет методы train() для обучения и toJSON() для получения внутреннего представления модели в формате JSON:

const brain = require('brain.js');
const net = new brain.NeuralNetwork();

net.train([
  { input: [0, 0], output: [0] },
  { input: [0, 1], output: [1] },
  { input: [1, 0], output: [1] },
  { input: [1, 1], output: [0] }
]);

const modelJSON = net.toJSON();

modelJSON содержит все необходимые параметры сети: веса, смещения, структуру слоев, функцию активации и настройки обучения.

Преобразование в ONNX

Для конвертации модели Brain.js в ONNX требуется промежуточная обработка. В текущей версии Brain.js прямой экспорт в ONNX не поддерживается, поэтому необходимо использовать вспомогательные библиотеки, такие как onnxjs или писать собственный скрипт конвертации.

  1. Извлечение структуры сети и весов

Каждый слой сети представлен объектом с указанием весов (weights), смещений (biases) и функции активации (activation). Для ONNX необходимо преобразовать эти данные в тензоры:

const onnx = require('onnxjs');

function convertToONNX(modelJSON) {
  const onnxMo del = {
    graph: {
      nodes: [],
      inputs: [],
      outputs: []
    }
  };

  modelJSON.layers.forEach((layer, index) => {
    const node = {
      name: `layer_${index}`,
      opType: 'Gemm',
      inputs: [/* данные весов */],
      outputs: [/* выход слоя */]
    };
    onnxModel.graph.nodes.push(node);
  });

  return onnxModel;
}
  1. Настройка входов и выходов модели

ONNX требует явного указания входных и выходных тензоров. Для нейронной сети с одним входом и одним выходом это выглядит следующим образом:

onnxModel.graph.inputs.push({
  name: 'input',
  type: 'float',
  shape: [1, modelJSON.inputSize]
});

onnxModel.graph.outputs.push({
  name: 'output',
  type: 'float',
  shape: [1, modelJSON.outputSize]
});
  1. Сохранение модели

После формирования структуры можно сохранить модель в формате .onnx с использованием fs или специализированных утилит:

const fs = require('fs');
fs.writeFileSync('model.onnx', JSON.stringify(onnxModel));

Для реального применения рекомендуется использовать библиотеку onnxjs для создания корректного бинарного файла, так как ONNX стандарт требует конкретного protobuf-формата.

Особенности при конвертации

  • Функции активации: Brain.js поддерживает sigmoid, relu, leaky-relu и tanh. При экспорте необходимо правильно сопоставить их с ONNX-операциями (Sigmoid, Relu, Tanh).
  • Слои: В Brain.js модель может состоять из произвольного количества скрытых слоев. ONNX представляется графом узлов, где каждый слой соответствует операции Gemm или MatMul + Add.
  • Точность и тип данных: Brain.js использует числа с плавающей запятой, обычно float32. ONNX также поддерживает float32, что упрощает конвертацию без потерь точности.
  • Обучение и инференс: ONNX ориентирован на инференс. Обученные веса из Brain.js можно использовать напрямую, но повторное обучение через ONNX не поддерживается без дополнительных средств.

Проверка модели

После конвертации рекомендуется протестировать ONNX-модель с помощью среды, поддерживающей ONNX (например, onnxruntime) и сравнить результаты с исходной моделью Brain.js, чтобы убедиться, что функциональность и точность сохранены.

const ort = require('onnxruntime-node');

async function testModel(inputData) {
  const session = await ort.InferenceSession.create('model.onnx');
  const feeds = { input: new ort.Tensor('float32', inputData, [1, inputData.length]) };
  const results = await session.run(feeds);
  console.log(results.output.data);
}

Тестирование позволяет выявить возможные несоответствия, например, из-за разницы в реализации функций активации или порядка слоев.

Практическое применение

Экспорт в ONNX позволяет использовать Brain.js-модели в:

  • Серверных приложениях на Python или C++ через ONNX Runtime.
  • Мобильных приложениях с поддержкой ONNX.
  • Системах с ускорением на GPU, где ONNX оптимизирован для высокопроизводительных вычислений.

Поддержка ONNX расширяет возможности интеграции, делая модели Brain.js переносимыми и совместимыми с промышленными решениями машинного обучения.