Инкрементальный инференс: KV-кэш и авторегрессионные модели

ONNX Runtime Web (ORT Web) представляет собой высокопроизводительную библиотеку для запуска моделей ONNX непосредственно в браузере или в среде Node.js. Основной задачей ORT Web является выполнение инференса с использованием различных бэкендов, таких как WebAssembly (WASM) и WebGPU, обеспечивая при этом переносимость моделей и минимальные накладные расходы на инфраструктуру.

Модель в ONNX загружается через InferenceSession, который является точкой входа для выполнения всех вычислений. Создание сессии требует указания пути к файлу модели или буфера с её содержимым. Базовый пример:

import * as ort from 'onnxruntime-web';

const session = await ort.InferenceSession.create('./model.onnx', {
  executionProviders: ['wasm', 'webgl'] // порядок приоритета
});

При инициализации можно задать настройки конкретного провайдера, например уровень параллелизма для WASM или выбор устройства для WebGPU.


KV-кэш и авторегрессионный инференс

Авторегрессионные модели, такие как GPT, используют предсказание токена за токеном. Каждый новый токен вычисляется на основе предыдущих, что при naïve реализации приводит к повторной обработке всей последовательности заново. KV-кэш (Key-Value cache) решает эту проблему, позволяя сохранять промежуточные представления key и value для каждого слоя трансформера, чтобы при генерации следующего токена повторно использовать уже вычисленные значения.

Структура KV-кэша

KV-кэш хранится в виде тензоров с размерами [num_heads, seq_len, head_dim] для каждого слоя. При каждом шаге инференса:

  1. Новый токен кодируется в эмбеддинг.
  2. Эмбеддинг проходит через слои трансформера, обновляя KV-кэш.
  3. В последующем шаге модель получает не всю историю, а только последний токен и текущий KV-кэш.

Пример создания KV-кэша в ORT Web:

// Инициализация KV-кэша для модели с N слоями
const kvCache = {};
for (let i = 0; i < numLayers; i++) {
  kvCache[`layer_${i}_key`] = new ort.Tensor('float32', new Float32Array(maxSeqLen * numHeads * headDim), [1, numHeads, maxSeqLen, headDim]);
  kvCache[`layer_${i}_value`] = new ort.Tensor('float32', new Float32Array(maxSeqLen * numHeads * headDim), [1, numHeads, maxSeqLen, headDim]);
}

Каждый шаг инференса теперь использует этот кэш:

const feeds = {
  input_ids: new ort.Tensor('int32', inputIds, [1, inputIds.length]),
  ...kvCache
};

const results = await session.run(feeds);

// Обновление KV-кэша
for (let i = 0; i < numLayers; i++) {
  kvCache[`layer_${i}_key`] = results[`layer_${i}_key`];
  kvCache[`layer_${i}_value`] = results[`layer_${i}_value`];
}

Управление последовательностью и генерацией

При работе с авторегрессионными моделями важна организация последовательности:

  • input_ids – массив токенов для текущего шага.
  • attention_mask – маска, указывающая на актуальные токены для внимания.
  • position_ids – идентификаторы позиций, необходимы для корректного использования слоев позиционного кодирования.

KV-кэш позволяет хранить позиционные представления каждого токена без необходимости повторного пересчёта всей последовательности. Это критично для длинных текстов и моделей с сотнями миллионов параметров.


Оптимизация инференса в браузере

ORT Web поддерживает несколько стратегий ускорения:

  1. WebAssembly (WASM) – кроссплатформенный вариант с низкой зависимостью от оборудования, но с ограниченной производительностью на GPU.
  2. WebGPU – доступ к видеоускорению через современные браузеры, позволяет использовать большие батчи и сокращать время вычислений KV-кэша.
  3. Батчинг – объединение нескольких запросов к модели в один пакет, минимизирует накладные расходы на запуск сессии и передачу данных.

Пример запуска с WebGPU:

const session = await ort.InferenceSession.create('./model.onnx', {
  executionProviders: ['webgpu'],
  graphOptimizationLevel: 'all'
});

Интеграция с генерацией текста

Алгоритм пошаговой генерации текста с KV-кэшем:

  1. Получение токена из текущего словаря.
  2. Формирование входного тензора с последним токеном.
  3. Прокрутка через модель с использованием KV-кэша.
  4. Обновление KV-кэша для следующего шага.
  5. Выбор следующего токена методом argmax или сэмплингом (top-k, top-p).

Эта схема позволяет поддерживать непрерывный контекст без повторного инференса всей последовательности.


Управление памятью и ресурсами

KV-кэш растёт пропорционально длине последовательности, поэтому важно:

  • Ограничивать maxSeqLen в модели и сессии.
  • Использовать циклическое обновление кэша для длинных текстов.
  • Освобождать неиспользуемые тензоры для предотвращения утечек памяти в браузере.

Пример очистки кэша:

for (const key in kvCache) {
  kvCache[key].data = null;
}

Это освобождает память, не разрушая структуру KV-кэша, чтобы его можно было инициализировать заново при следующем инференсе.


Особенности ONNX Runtime Web для авторегрессионных моделей

  • Поддержка мульти-выходных моделей, где каждый слой возвращает ключи и значения для KV-кэша.
  • Лёгкая интеграция с фронтенд-приложениями без необходимости серверного инференса.
  • Возможность использования любых моделей ONNX, экспортированных из PyTorch или TensorFlow, с сохранением структуры KV-кэша и авторегрессионной логики.

Ключевой принцип работы — хранение промежуточных представлений и минимизация повторных вычислений для каждого шага генерации. Такой подход делает ORT Web эффективным инструментом для интерактивного и потокового инференса моделей трансформеров.