Постобработка классификации: softmax, argmax, top-k

ONNX Runtime Web (ORT Web) предоставляет возможность запускать модели машинного обучения непосредственно в браузере с использованием JavaScript. После получения выходных данных модели классификации важным этапом является их корректная постобработка. Наиболее часто применяемыми методами являются softmax, argmax и top-k, обеспечивающие интерпретацию сырых логитов в удобочитаемые вероятности и метки классов.


Softmax

Softmax — функция активации, которая преобразует набор сырых выходов модели (логитов) в распределение вероятностей. Для каждого элемента вектора ( z_i ) функция вычисляется по формуле:

[ (z_i) = ]

Особенности применения:

  • Преобразует любые вещественные значения логитов в диапазон ([0, 1]), суммарно равный 1.
  • Позволяет интерпретировать результат как вероятность принадлежности к конкретному классу.
  • В ONNX Runtime Web реализуется через обычную математику Jav * aScript: экспоненту, суммирование и деление.

Пример реализации softmax на Jav * aScript:

function softmax(logits) {
    const maxLogit = Math.max(...logits);
    const exps = logits.map(x => Math.exp(x - maxLogit));
    const sumExps = exps.reduce((a, b) => a + b, 0);
    return exps.map(e => e / sumExps);
}

Сдвиг на maxLogit предотвращает переполнение экспоненты при больших значениях логитов, что является стандартной практикой численной стабильности.


Argmax

Argmax используется для определения индекса класса с наибольшей вероятностью. После применения softmax или непосредственно к логитам, argmax позволяет получить итоговую категорию.

Применение:

  • Вектор вероятностей softmax используется для выбора класса с максимальной вероятностью.
  • В ONNX Runtime Web вычисляется простым перебором массива и сравнением элементов.

Пример реализации argmax на Jav * aScript:

function argmax(array) {
    let maxIndex = 0;
    let maxValue = array[0];
    for (let i = 1; i < array.length; i++) {
        if (array[i] > maxValue) {
            maxValue = array[i];
            maxIndex = i;
        }
    }
    return maxIndex;
}

Это обеспечивает точное определение наиболее вероятного класса без использования сторонних библиотек.


Top-K

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

Особенности:

  • K — число возвращаемых наиболее вероятных классов.
  • Возвращается массив объектов, содержащий индекс класса и его вероятность.
  • Реализуется через сортировку и срез массива.

Пример реализации top-k на Jav * aScript:

function topK(array, k) {
    return array
        .map((value, index) => ({ index, value }))
        .sort((a, b) => b.value - a.value)
        .slice(0, k);
}

// Пример использования:
const probabilities = softmax([2.5, 0.3, 2.1, 0.7]);
const top3 = topK(probabilities, 3);
console.log(top3);

Особенности производительности:

  • Для больших векторов логитов сортировка всех элементов может быть неэффективной.
  • В задачах с высоким числом классов рекомендуется использовать алгоритмы частичной сортировки, чтобы получить top-k без полной сортировки.

Интеграция с ONNX Runtime Web

После запуска модели в ORT Web через session.run() полученные выходные данные представляют собой TypedArray. Для постобработки необходимо:

  1. Преобразовать TypedArray в обычный JavaScript-массив.
  2. Применить softmax для получения вероятностей.
  3. Использовать argmax или top-k для извлечения наиболее вероятных классов.

Пример полного пайплайна:

const session = await ort.InferenceSession.create('model.onnx');
const feeds = { input: new ort.Tensor('float32', inputData, [1, inputSize]) };
const results = await session.run(feeds);
const logits = Array.from(results.output.data);

const probabilities = softmax(logits);
const top5 = topK(probabilities, 5);

console.log(top5);

Этот подход обеспечивает точную и прозрачную интерпретацию результатов модели прямо в браузере, без серверной обработки.


Рекомендации по точности

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

Такая комбинация softmax, argmax и top-k формирует стандартный и эффективный пайплайн постобработки классификационных моделей в ONNX Runtime Web, позволяя корректно интерпретировать результаты и извлекать релевантные классы для дальнейшего использования.