Генерация текста: GPT-2, небольшие языковые модели

ONNX Runtime Web (ORT Web) предоставляет высокопроизводительное выполнение моделей ONNX прямо в браузере на JavaScript. Для начала работы требуется установка пакета через npm:

npm install onnxruntime-web

После установки библиотеку можно импортировать в проект:

import * as ort from 'onnxruntime-web';

ORT Web поддерживает несколько бэкендов: webgl для ускорения на GPU через WebGL и wasm для выполнения через WebAssembly на CPU. Выбор бэкенда влияет на производительность:

const session = await ort.InferenceSession.create('gpt2.onnx', { executionProviders: ['webgl'] });

Загрузка и подготовка модели GPT-2

Модель GPT-2 в формате ONNX обычно включает два основных компонента: энкодер токенов и саму нейросеть. Для работы требуется файл модели gpt2.onnx и соответствующий токенизатор.

Токенизация выполняется через библиотеку @huggingface/tokenizers или вручную, если используется простая схема:

import { GPT2Tokenizer } from '@huggingface/tokenizers';

const tokenizer = await GPT2Tokenizer.fromOptions({ model: 'gpt2' });
const inputIds = tokenizer.encode('Пример текста для генерации').ids;

Полученные токены преобразуются в тензор, подходящий для ONNX Runtime:

const inputTensor = new ort.Tensor('int64', BigInt64Array.from(inputIds.map(BigInt)), [1, inputIds.length]);

Создание сессии и выполнение инференса

После подготовки модели и входных данных создаётся сессия инференса. Важный момент — правильное указание формы входного тензора. Для GPT-2 она имеет вид [batch_size, sequence_length]. Для однократного запуска используется batch_size = 1.

const feeds = { input_ids: inputTensor };
const results = await session.run(feeds);
const outputIds = results['output_ids'].data;

Если модель поддерживает past_key_values, их можно передавать для ускорения генерации длинного текста:

let pastKeyValues = null;

const feeds = { input_ids: inputTensor };
if (pastKeyValues) {
    Object.assign(feeds, pastKeyValues);
}

const results = await session.run(feeds);
pastKeyValues = extractPastKeyValues(results);

Обработка и декодирование выхода

Результат инференса GPT-2 — массив индексов слов из словаря токенизатора. Для получения читаемого текста необходимо декодирование:

const generatedText = tokenizer.decode(Array.from(outputIds), { skipSpecialTokens: true });

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

Настройка параметров генерации

Параметры генерации влияют на качество и разнообразие текста:

  • max_length — максимальная длина генерируемой последовательности.
  • temperature — регулировка вероятности выбора токенов; низкие значения делают текст более предсказуемым.
  • top_k — ограничение выбора токенов по наивысшей вероятности.
  • top_p — ядерная выборка: сумма вероятностей токенов ограничена до значения p.

Пример генерации с контролем параметров:

async function generateText(session, tokenizer, prompt, maxLength = 50, temperature = 0.7, topK = 50, topP = 0.9) {
    let inputIds = tokenizer.encode(prompt).ids;
    let pastKeyValues = null;

    for (let step = 0; step < maxLength; step++) {
        const inputTensor = new ort.Tensor('int64', BigInt64Array.from(inputIds.map(BigInt)), [1, inputIds.length]);
        const feeds = { input_ids: inputTensor };
        if (pastKeyValues) Object.assign(feeds, pastKeyValues);

        const results = await session.run(feeds);
        pastKeyValues = extractPastKeyValues(results);

        const logits = results['logits'].data;
        const nextToken = sampleNextToken(logits, temperature, topK, topP);

        inputIds.push(nextToken);
    }

    return tokenizer.decode(inputIds, { skipSpecialTokens: true });
}

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

  • Использование webgl ускоряет вычисления на GPU, особенно для длинных последовательностей.
  • Поддержка past_key_values значительно снижает количество операций для последовательного добавления токенов.
  • В браузере стоит следить за потреблением памяти, очищать тензоры, которые больше не нужны, чтобы избежать утечек.

Работа с несколькими моделями

Можно одновременно загружать несколько моделей GPT-2 различной мощности (small, medium, large). Для этого создаются отдельные сессии:

const smallSession = await ort.InferenceSession.create('gpt2-small.onnx', { executionProviders: ['webgl'] });
const mediumSession = await ort.InferenceSession.create('gpt2-medium.onnx', { executionProviders: ['webgl'] });

Использование подходящей модели зависит от задачи: маленькие модели быстрее и требуют меньше памяти, большие — генерируют более связный и богатый текст.

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

Для одновременной генерации нескольких текстов удобно использовать батчинг, формируя входной тензор с формой [batch_size, sequence_length]. ORT Web корректно обрабатывает многомерные тензоры, позволяя ускорять обработку при параллельном генеративном инференсе.