Multi-Head Attention является ключевым компонентом современных архитектур нейронных сетей, включая трансформеры. Он позволяет модели одновременно учитывать разные аспекты входных данных, обеспечивая более богатое представление контекстной информации. В TensorFlow.js реализация этого механизма строится на базовых операциях с тензорами и слоях Keras.
Multi-Head Attention расширяет концепцию Scaled Dot-Product Attention, выполняя несколько параллельных вычислений внимания (heads) и объединяя их результаты. Формально процесс можно представить так:
[ (Q, K, V) = () V]
где (d_k) — размерность ключей.
Эта схема позволяет модели захватывать различные типы зависимостей в данных параллельно.
TensorFlow.js предоставляет возможности для создания многоуровневой
архитектуры внимания с помощью API tf.layers и операций с
тензорами.
Простейший способ реализации — вручную определить веса для 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:
В 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 без серверной поддержки.
Этот подход позволяет создавать высокопроизводительные модели внимания на JavaScript с возможностью интерактивной работы в браузере, обучая и применяя трансформеры без необходимости серверной инфраструктуры.