Трансформеры с attention mask переменной длины

ONNX Runtime Web (ORT Web) представляет собой высокопроизводительную библиотеку для выполнения моделей машинного обучения прямо в браузере с поддержкой WebAssembly и WebGPU. При работе с трансформерами, особенно моделями с механизмом внимания (attention), ключевым моментом является управление масками внимания переменной длины.

Подключение и инициализация ORT Web

Для работы с ORT Web необходимо установить пакет через npm:

import * as ort from 'onnxruntime-web';

// Выбор движка: 'wasm' или 'webgl' (WebGPU)
const sessionOptions = {
  executionProviders: ['wasm'], // можно использовать 'webgl' или 'webgpu' в зависимости от поддержки
  graphOptimizationLevel: 'all',
};

const session = await ort.InferenceSession.create('model.onnx', sessionOptions);

Ключевые параметры:

  • executionProviders – список доступных провайдеров вычислений.
  • graphOptimizationLevel – оптимизация графа модели для ускорения инференса.

Структура входных данных для трансформеров

Трансформеры используют три основных входа:

  1. input_ids – тензор с индексами токенов.
  2. attention_mask – тензор, указывающий, какие токены должны быть видимы для механизма внимания.
  3. token_type_ids (опционально) – для моделей типа BERT, определяет сегменты текста.

Формат тензора для ORT Web:

const inputTensor = new ort.Tensor('int64', inputIdsArray, [batchSize, sequenceLength]);
const attentionTensor = new ort.Tensor('int64', attentionMaskArray, [batchSize, sequenceLength]);

Использование attention mask переменной длины

В трансформерах attention mask служит для игнорирования паддинговых токенов при вычислении внимания. В реальных сценариях длина входной последовательности может меняться от батча к батчу.

Проблема: стандартная ONNX-модель ожидает фиксированную форму тензора. Чтобы поддерживать переменную длину:

  • Использовать динамический sequenceLength при экспорте модели в ONNX (dynamic_axes).
  • Перед инференсом правильно формировать attention mask под текущую длину последовательности.

Пример создания attention mask:

function createAttentionMask(inputIds, padTokenId = 0) {
  return inputIds.map(seq => seq.map(token => token === padTokenId ? 0 : 1));
}

const attentionMaskArray = createAttentionMask(inputIdsArray);
const attentionMaskTensor = new ort.Tensor('int64', attentionMaskArray.flat(), [batchSize, sequenceLength]);

Инференс с ORT Web

const feeds = {
  input_ids: inputTensor,
  attention_mask: attentionMaskTensor
};

const results = await session.run(feeds);
const output = results['last_hidden_state']; // пример для BERT-подобной модели

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

  • Входные данные и маска должны иметь одинаковую размерность [batchSize, sequenceLength].
  • Если последовательность короче, необходимо дополнять паддингом и корректно формировать маску.

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

  1. Выбор движка: WebGPU обеспечивает максимальную производительность на современных устройствах, WebAssembly – на старых браузерах.
  2. Оптимизация графа: graphOptimizationLevel: 'all' минимизирует избыточные вычисления.
  3. Минимизация конверсий: хранение входных массивов в типе TypedArray уменьшает накладные расходы на копирование данных.

Динамическая длина батчей

Для реальных приложений часто требуется батчинг с переменной длиной последовательностей. Стратегия:

  • Сортировка последовательностей по длине перед формированием батча.
  • Паддинг до максимальной длины в батче.
  • Генерация attention mask с учетом фактической длины.

Пример формирования батча:

function padBatch(sequences, padTokenId = 0) {
  const maxLen = Math.max(...sequences.map(seq => seq.length));
  const padded = sequences.map(seq => [...seq, ...Array(maxLen - seq.length).fill(padTokenId)]);
  const mask = sequences.map(seq => [...Array(seq.length).fill(1), ...Array(maxLen - seq.length).fill(0)]);
  return { padded, mask };
}

const { padded, mask } = padBatch(inputIdsArray);

Совместимость с Hugging Face и ONNX

Модели, экспортированные из Hugging Face, часто поддерживают динамическую длину входа через параметр dynamic_axes. Для корректного использования в ORT Web:

  • Проверить, что sequence_length и batch_size являются динамическими.
  • При конвертации модели из PyTorch использовать torch.onnx.export с соответствующими динамическими осями:
torch.onnx.export(
    model,
    (dummy_input, dummy_attention),
    "model.onnx",
    input_names=["input_ids", "attention_mask"],
    output_names=["last_hidden_state"],
    dynamic_axes={
        "input_ids": {0: "batch_size", 1: "sequence_length"},
        "attention_mask": {0: "batch_size", 1: "sequence_length"},
        "last_hidden_state": {0: "batch_size", 1: "sequence_length"}
    },
)

Работа с многомерными масками

Для моделей типа GPT-2 или T5 attention mask может иметь форму [batch_size, num_heads, sequence_length, sequence_length]. В ORT Web необходимо правильно создавать такие тензоры:

const numHeads = 12;
const seqLen = 16;
const mask = new Array(batchSize * numHeads * seqLen * seqLen).fill(0);

// пример заполнения mask: 1 там, где токен видим

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