ONNX Runtime Web (ORT Web) представляет собой высокопроизводительную библиотеку для выполнения моделей машинного обучения прямо в браузере с поддержкой WebAssembly и WebGPU. При работе с трансформерами, особенно моделями с механизмом внимания (attention), ключевым моментом является управление масками внимания переменной длины.
Для работы с 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 – оптимизация графа модели для
ускорения инференса.Трансформеры используют три основных входа:
Формат тензора для ORT Web:
const inputTensor = new ort.Tensor('int64', inputIdsArray, [batchSize, sequenceLength]);
const attentionTensor = new ort.Tensor('int64', attentionMaskArray, [batchSize, sequenceLength]);
В трансформерах attention mask служит для игнорирования паддинговых токенов при вычислении внимания. В реальных сценариях длина входной последовательности может меняться от батча к батчу.
Проблема: стандартная ONNX-модель ожидает фиксированную форму тензора. Чтобы поддерживать переменную длину:
sequenceLength при экспорте
модели в ONNX (dynamic_axes).Пример создания 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]);
const feeds = {
input_ids: inputTensor,
attention_mask: attentionMaskTensor
};
const results = await session.run(feeds);
const output = results['last_hidden_state']; // пример для BERT-подобной модели
Особенности:
[batchSize, sequenceLength].graphOptimizationLevel: 'all' минимизирует избыточные
вычисления.TypedArray уменьшает накладные расходы на копирование
данных.Для реальных приложений часто требуется батчинг с переменной длиной последовательностей. Стратегия:
Пример формирования батча:
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, часто поддерживают
динамическую длину входа через параметр dynamic_axes. Для
корректного использования в ORT Web:
sequence_length и
batch_size являются динамическими.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 там, где токен видим
Правильная структура маски обеспечивает корректное применение внимания, особенно при генерации текста с динамической длиной контекста.