Постобработка генерации токенов и beam search в JS

Для работы с ONNX Runtime Web используется пакет onnxruntime-web, который позволяет запускать модели ONNX прямо в браузере или в среде Node.js с WebAssembly или WebGL. Установка выполняется через npm:

npm install onnxruntime-web

Инициализация окружения включает импорт библиотеки и создание сессии для модели:

import * as ort from 'onnxruntime-web';

const session = await ort.InferenceSession.create('model.onnx', {
  executionProviders: ['wasm'] // 'webgl' для ускорения через GPU
});

Сессия отвечает за загрузку модели и управление вычислительными графами. executionProviders позволяет выбирать между CPU (WASM) и GPU (WebGL), что влияет на производительность.


Формирование входных данных и типы тензоров

ONNX Runtime Web использует ort.Tensor для передачи данных в модель. Тензоры должны соответствовать формату, ожидаемому моделью: тип данных (float32, int32 и др.) и размерность.

const inputIds = new ort.Tensor('int32', [101, 102, 103], [1, 3]); // shape [batch, sequence]
const attentionMask = new ort.Tensor('int32', [1, 1, 1], [1, 3]);
const inputs = { input_ids: inputIds, attention_mask: attentionMask };

Важно корректно задавать shape, так как несоответствие может привести к ошибкам выполнения.


Постобработка генерации токенов

После получения сырых выходов модели (logits) необходимо провести несколько ключевых шагов для превращения чисел в текст:

  1. Применение софтмакса Модель возвращает логиты — необработанные оценки вероятностей для каждого токена. Для получения вероятностей используется функция softmax:
function softmax(logits) {
  const maxLogit = Math.max(...logits);
  const exps = logits.map(l => Math.exp(l - maxLogit));
  const sumExps = exps.reduce((a, b) => a + b, 0);
  return exps.map(e => e / sumExps);
}
  1. Выбор токена После преобразования логитов в вероятности можно выбрать следующий токен. Наиболее простой подход — greedy decoding:
const nextTokenIndex = probabilities.indexOf(Math.max(...probabilities));

Однако greedy decoding часто приводит к менее разнообразным и повторяющимся текстам.

  1. Применение топ-k и топ-p фильтров Для генерации более разнообразного текста используются методы с ограничением вероятностей:
function topKSampling(probabilities, k) {
  const sorted = probabilities
    .map((p, i) => [i, p])
    .sort((a, b) => b[1] - a[1])
    .slice(0, k);
  const sum = sorted.reduce((a, b) => a + b[1], 0);
  const normalized = sorted.map(([i, p]) => [i, p / sum]);
  const r = Math.random();
  let cumulative = 0;
  for (const [i, p] of normalized) {
    cumulative += p;
    if (r < cumulative) return i;
  }
}

Beam Search: поиск с ветвлением

Beam search позволяет учитывать несколько кандидатов на каждом шаге, улучшая качество последовательности. Основные параметры:

  • beamWidth — количество одновременно отслеживаемых гипотез;
  • lengthPenalty — штраф за короткие последовательности;
  • eosTokenId — токен завершения последовательности.

Структура алгоритма

  1. Инициализация beams:
let beams = [{ tokens: [startToken], score: 0 }];
  1. Расширение каждого бима:
function expandBeams(beams, logits, beamWidth) {
  const allCandidates = [];
  for (const beam of beams) {
    const probabilities = softmax(logits[beam.tokens[beam.tokens.length - 1]]);
    const topIndices = probabilities
      .map((p, i) => [i, p])
      .sort((a, b) => b[1] - a[1])
      .slice(0, beamWidth);
    for (const [token, prob] of topIndices) {
      allCandidates.push({
        tokens: [...beam.tokens, token],
        score: beam.score + Math.log(prob)
      });
    }
  }
  return allCandidates
    .sort((a, b) => b.score - a.score)
    .slice(0, beamWidth);
}
  1. Проверка на завершение последовательности:
const completed = beams.filter(b => b.tokens[b.tokens.length - 1] === eosTokenId);

Beam search позволяет сохранять несколько высоко вероятных вариантов текста и выбирать наиболее осмысленные результаты.


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

  • Использование WebGL вместо WASM значительно ускоряет вычисления на больших моделях.
  • Минимизация конверсий между массивами JavaScript и ort.Tensor снижает накладные расходы.
  • Важно использовать одну сессию для нескольких шагов генерации, чтобы избежать повторной загрузки модели.

let beams = [{ tokens: [startToken], score: 0 }];
for (let step = 0; step < maxLength; step++) {
  const inputTensor = new ort.Tensor('int32', beams.map(b => b.tokens), [beams.length, step + 1]);
  const outputs = await session.run({ input_ids: inputTensor });
  beams = expandBeams(beams, outputs.logits, beamWidth);
  if (beams.every(b => b.tokens[b.tokens.length - 1] === eosTokenId)) break;
}

const bestSequence = beams[0].tokens;

Этот подход объединяет строгий контроль вероятностей на каждом шаге и сохранение нескольких кандидатов, обеспечивая высокое качество генерации текста.


Если требуется, можно рассмотреть дополнительные методы постобработки, такие как повторное взвешивание вероятностей с учётом частот токенов, наложение ограничений по контексту или фильтрация стоп-слов. Все эти методы органично интегрируются с ONNX Runtime Web через работу с logits и массивами токенов на JavaScript.