Верификация модели путём сравнения с эталонными выходами

ONNX Runtime Web (ORT Web) предоставляет эффективный способ запуска моделей машинного обучения в браузере, используя JavaScript. Ключевой аспект работы с любыми моделями — проверка корректности их работы. Верификация модели основывается на сравнении её предсказаний с заранее подготовленными эталонными выходами (reference outputs). Это обеспечивает уверенность в том, что модель загружена и выполняется корректно, а её вычисления соответствуют ожидаемым результатам.


Подготовка модели и эталонных данных

Для начала требуется подготовить следующие компоненты:

  1. ONNX-модель, экспортированная из фреймворка (PyTorch, TensorFlow и др.).
  2. Эталонные входные данные (input tensors), соответствующие формату модели.
  3. Эталонные выходные данные (expected outputs), полученные с помощью проверенной реализации или контрольного запуска модели.

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

import * as ort from "onnxruntime-web";

// Пример загрузки эталонных входных данных
const referenceInputs = await fetch("reference_inputs.json").then(res => res.json());
const referenceOutputs = await fetch("reference_outputs.json").then(res => res.json());

Инициализация среды ONNX Runtime Web

ORT Web поддерживает несколько бэкендов: WebAssembly (wasm), WebGL (webgl) и WebGPU (webgpu). Для верификации обычно используется wasm для максимальной предсказуемости вычислений, так как аппаратные ускорители могут давать незначительные различия из-за особенностей чисел с плавающей точкой.

const session = await ort.InferenceSession.create("model.onnx", {
  executionProviders: ["wasm"] // Гарантированная повторяемость вычислений
});

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

Входные данные должны соответствовать именам и типам входов модели. В ORT Web используется объект, где ключи — это имена входных тензоров, а значения — объекты ort.Tensor.

const inputs = {};
for (const name in referenceInputs) {
  const data = new Float32Array(referenceInputs[name].data);
  const dims = referenceInputs[name].dims;
  inputs[name] = new ort.Tensor("float32", data, dims);
}

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

Сессия ONNX Runtime Web позволяет вызвать метод run, передав подготовленные входы. Метод возвращает объект, где ключи — это имена выходных тензоров, а значения — объекты ort.Tensor.

const results = await session.run(inputs);

Сравнение с эталонными выходами

Для проверки корректности необходимо сравнивать значения предсказаний с эталонными значениями. Поскольку вычисления с плавающей точкой могут иметь малые погрешности, применяется сравнение с допуском (tolerance).

function arraysAreClose(arr1, arr2, tol = 1e-5) {
  if (arr1.length !== arr2.length) return false;
  for (let i = 0; i < arr1.length; i++) {
    if (Math.abs(arr1[i] - arr2[i]) > tol) return false;
  }
  return true;
}

for (const name in referenceOutputs) {
  const predicted = results[name].data;
  const expected = new Float32Array(referenceOutputs[name].data);
  if (!arraysAreClose(predicted, expected)) {
    console.error(`Несовпадение в выходе тензора ${name}`);
  }
}

Ключевые моменты:

  • Использование Float32Array гарантирует правильное сравнение чисел с плавающей точкой.
  • Значение tol должно подбираться с учётом погрешностей вычислений на выбранном бэкенде.
  • При обнаружении несовпадений важно проверять корректность как входных данных, так и порядка элементов в тензоре.

Автоматизация верификации

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

async function verifyModel(session, testCases, tolerance = 1e-5) {
  let successCount = 0;
  for (const testCase of testCases) {
    const inputs = {};
    for (const name in testCase.inputs) {
      const { data, dims } = testCase.inputs[name];
      inputs[name] = new ort.Tensor("float32", new Float32Array(data), dims);
    }
    const results = await session.run(inputs);
    let allClose = true;
    for (const name in testCase.outputs) {
      const predicted = results[name].data;
      const expected = new Float32Array(testCase.outputs[name].data);
      if (!arraysAreClose(predicted, expected, tolerance)) {
        console.warn(`Несовпадение в ${name}`);
        allClose = false;
      }
    }
    if (allClose) successCount++;
  }
  console.log(`Верификация завершена: ${successCount}/${testCases.length} тестов успешны`);
}

Особенности числовой точности и аппаратных ускорителей

  • WebAssembly (wasm): наиболее стабильный и предсказуемый результат.
  • WebGL и WebGPU: могут давать небольшие расхождения в последней значащей цифре из-за оптимизаций графического процессора.
  • Допуск сравнения (tolerance) следует выбирать исходя из модели и используемого бэкенда. Для классификационных задач можно использовать 1e-5–1e-4, для регрессионных — 1e-6–1e-7.

Вывод

Процесс верификации модели в ONNX Runtime Web строится на последовательных шагах: загрузка модели, подготовка входов и эталонных данных, выполнение модели и сравнение результатов с эталонными значениями. Такой подход позволяет выявить ошибки в экспорте модели, некорректную обработку входных данных и проблемы с численной точностью, обеспечивая надежную работу модели в веб-приложении.