Multi-head attention

Multi-Head Attention является ключевым компонентом современных архитектур нейронных сетей, включая трансформеры. Он позволяет модели одновременно учитывать разные аспекты входных данных, обеспечивая более богатое представление контекстной информации. В TensorFlow.js реализация этого механизма строится на базовых операциях с тензорами и слоях Keras.


Основная идея

Multi-Head Attention расширяет концепцию Scaled Dot-Product Attention, выполняя несколько параллельных вычислений внимания (heads) и объединяя их результаты. Формально процесс можно представить так:

  1. Для каждого “head” создаются проекции входных матриц Query (Q), Key (K) и Value (V) с помощью обучаемых весов.
  2. Вычисляется внимание через скалярное произведение:

[ (Q, K, V) = () V]

где (d_k) — размерность ключей.

  1. Результаты всех heads объединяются и проходят через финальный линейный слой.

Эта схема позволяет модели захватывать различные типы зависимостей в данных параллельно.


Реализация в TensorFlow.js

TensorFlow.js предоставляет возможности для создания многоуровневой архитектуры внимания с помощью API tf.layers и операций с тензорами.

Создание слоя Multi-Head Attention

Простейший способ реализации — вручную определить веса для Q, K и V, затем вычислить attention и объединить heads.

import * as tf from '@tensorflow/tfjs';

class MultiHeadAttention extends tf.layers.Layer {
  constructor(config) {
    super(config);
    this.numHeads = config.numHeads;
    this.modelDim = config.modelDim;
    this.keyDim = this.modelDim / this.numHeads;
  }

  build(inputShape) {
    this.Wq = this.addWeight('Wq', [this.modelDim, this.modelDim], 'float32', tf.initializers.glorotUniform());
    this.Wk = this.addWeight('Wk', [this.modelDim, this.modelDim], 'float32', tf.initializers.glorotUniform());
    this.Wv = this.addWeight('Wv', [this.modelDim, this.modelDim], 'float32', tf.initializers.glorotUniform());
    this.Wo = this.addWeight('Wo', [this.modelDim, this.modelDim], 'float32', tf.initializers.glorotUniform());
  }

  call(inputs) {
    const [query, key, value] = inputs;

    // Линейные проекции
    let Q = tf.matMul(query, this.Wq.read());
    let K = tf.matMul(key, this.Wk.read());
    let V = tf.matMul(value, this.Wv.read());

    // Разделение на heads
    Q = this.splitHeads(Q);
    K = this.splitHeads(K);
    V = this.splitHeads(V);

    // Scaled dot-product attention
    let attentionOutput = this.scaledDotProductAttention(Q, K, V);

    // Объединение heads
    attentionOutput = this.combineHeads(attentionOutput);

    return tf.matMul(attentionOutput, this.Wo.read());
  }

  splitHeads(x) {
    const [batchSize, seqLen, dim] = x.shape;
    return tf.reshape(x, [batchSize, seqLen, this.numHeads, this.keyDim])
             .transpose([0, 2, 1, 3]);
  }

  combineHeads(x) {
    const [batchSize, numHeads, seqLen, depth] = x.shape;
    return tf.reshape(x.transpose([0, 2, 1, 3]), [batchSize, seqLen, numHeads * depth]);
  }

  scaledDotProductAttention(Q, K, V) {
    const matmulQK = tf.matMul(Q, K, false, true);
    const dk = tf.scalar(this.keyDim, 'float32');
    const scaled = tf.div(matmulQK, tf.sqrt(dk));
    const weights = tf.softmax(scaled);
    return tf.matMul(weights, V);
  }
}

Ключевые моменты реализации

  • Разделение и объединение heads Для эффективной работы модели необходимо правильно транспонировать и ресайзить тензоры, чтобы каждая голова обрабатывала отдельный подпространственный сегмент.

  • Масштабирование dot-product Деление на () предотвращает слишком большие значения перед softmax, что стабилизирует обучение.

  • Обучаемые веса Каждый head имеет свои проекции Q, K, V, а также общий выходной линейный слой, что позволяет модели гибко комбинировать различные аспекты внимания.


Использование в архитектуре трансформера

Multi-Head Attention может применяться как в self-attention, так и в encoder-decoder attention:

  • Self-Attention: Q, K и V берутся из одного и того же входа, что позволяет захватывать зависимости внутри последовательности.
  • Cross-Attention: Q формируется из декодера, а K и V — из энкодера, обеспечивая интеграцию информации между последовательностями.

В TensorFlow.js это реализуется аналогично, меняется лишь источник тензоров для Q, K и V.


Примеры интеграции в модель

const input = tf.input({shape: [seqLen, modelDim]});
const mhaLayer = new MultiHeadAttention({numHeads: 8, modelDim: modelDim});
const output = mhaLayer.apply([input, input, input]);
const model = tf.model({inputs: input, outputs: output});

Такой подход позволяет строить сложные трансформерные архитектуры полностью на стороне клиента в браузере, используя TensorFlow.js без серверной поддержки.


Практические рекомендации

  1. Подбор числа heads: оптимальное количество голов зависит от размерности модели; слишком много heads при малой размерности может привести к избыточной сложности и переобучению.
  2. Регуляризация: рекомендуется использовать dropout на attention weights и после линейного слоя Wo для улучшения обобщающей способности.
  3. Производительность: вычисления attention требуют значительных ресурсов, особенно на больших последовательностях; эффективнее использовать батчи и WebGL ускорение в TensorFlow.js.

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