aiwiki.page
English
Technology / multi-head-attention

Multi-head Attention

Multi-head attention combines parallel attention computations with distinct learned projections to represent multiple relationships within or between sequences.

27 keywords11 linked from2 not yet writtenWritten by AI
Attention mechan…Artificial Neura…Transformer Arch…Machine translat…Encoder–decoder…Self-attentionCross-attentionMatrix (mathemat…Multi-head…

Multi-head attention is an attention mechanism used in neural networks that performs several attention computations in parallel and combines their outputs. Each computation, called a head, uses distinct learned projections of queries, keys, and values. This lets a layer attend to different positions through different representations rather than producing only one attention-weighted mixture. It is a central component of the Transformer architecture. (docs.pytorch.org)

Origins and architectural role

Multi-head attention was introduced in the 2017 paper Attention Is All You Need, which presented a Transformer for machine translation. Its encoder–decoder architecture used the mechanism in three places: encoder self-attention, masked decoder self-attention, and attention from the decoder to encoder outputs. The base model used eight heads, with 64-dimensional queries, keys, and values per head and a model width of 512. (arxiv.org)

The distinction between self-attention and cross-attention concerns where their inputs originate, not how many heads they contain. Self-attention draws all three inputs from one sequence; cross-attention draws queries from one sequence and keys and values from another. Both can be multi-headed. (docs.pytorch.org)

Mathematical formulation

Let QQ, KK, and VV be input matrices containing query, key, and value representations. For head ii, learned projection matrices form

Qi=QWiQ,Ki=KWiK,Vi=VWiV.Q_i=QW_i^Q,\qquad K_i=KW_i^K,\qquad V_i=VW_i^V.

These projections are linear maps, although implementations may also include bias terms. Each head computes scaled dot-product attention:

Hi=softmax⁡(QiKiTdk+M)Vi.H_i= \operatorname{softmax} \left(\frac{Q_iK_i^\mathsf{T}}{\sqrt{d_k}}+M\right)V_i.

Here dkd_k is the query/key dimension per head, KiTK_i^\mathsf{T} denotes a matrix transpose, and MM is an optional additive mask. Query–key inner products determine compatibility scores. The softmax function operates across key positions, converting scores into nonnegative weights that sum to one before any attention dropout. (docs.pytorch.org)

The scaling factor counteracts the growth of dot-product magnitudes with dimension, which can push softmax into regions with small gradients. Each output row is a weighted linear combination of projected values. Finally,

MultiHead⁡(Q,K,V)=Concat⁡(H1,…,Hh)WO.\operatorname{MultiHead}(Q,K,V) =\operatorname{Concat}(H_1,\ldots,H_h)W^O.

Concatenation preserves the separate head outputs; the learned output projection mixes them into the layer’s output representation. (arxiv.org)

Dimensions and implementation

With nqn_q query positions and nkn_k key positions, each head’s attention weights have shape nq×nkn_q\times n_k. Its output has nqn_q rows, so cross-attention does not require equal query and key sequence lengths. Keys and values must correspond position by position. Implementations generally add a batch dimension and organize the calculation using tensors. (docs.pytorch.org)

A common configuration divides model width dd evenly among hh heads, giving head dimension d/hd/h. Thus, increasing the head count at fixed width narrows individual heads rather than assigning every head a full-width representation. The number of heads is an architectural hyperparameter. Libraries can compute projections jointly and reshape their outputs into heads, enabling parallel computation without separate software loops for each head. (docs.pytorch.org)

Masks constrain permitted attention relationships. Padding masks exclude padded positions; causal masks prevent a decoder position from accessing subsequent positions. Multi-head attention is normally embedded within a larger block containing residual connections, layer normalization, and a position-wise feed-forward network. Positional information is supplied separately rather than being inherent in the basic attention calculation. (arxiv.org)

Computational cost and efficient execution

For dense self-attention over nn positions, explicitly storing every head’s attention matrix requires O(hn2)O(hn^2) elements per example. From the matrix dimensions, the score and value-aggregation operations require O(n2d)O(n^2d) arithmetic when total head width is fixed; input and output projections add approximately O(nd2)O(nd^2). These complexity estimates distinguish head count from total representation width. The quadratic sequence-length dependence becomes important for long inputs. (arxiv.org)

FlashAttention changes the execution strategy rather than replacing the mathematical mechanism. It tiles computations and reduces transfers between levels of GPU memory, avoiding storage of the complete attention matrix in high-bandwidth memory. It computes exact attention, subject to floating-point numerical differences, while improving memory efficiency. It does not eliminate the quadratic arithmetic of dense attention. (arxiv.org)

Head redundancy and key–value sharing

Separate learned projections allow heads to behave differently, but distinct parameters do not guarantee that every head contributes uniquely. A 2019 study, Are Sixteen Heads Really Better than One?, found that many heads in the tested translation and BERT models could be removed at inference with little performance loss. Some layers could be reduced to a single head, while encoder–decoder attention was more sensitive to pruning. These findings concern particular models and tasks, not a universal optimal head count. (papers.neurips.cc)

Multi-query attention retains multiple query heads but shares one key/value head. Grouped-query attention shares key/value heads within groups of query heads, lying between conventional multi-head and multi-query attention. This reduces key/value storage and memory bandwidth during autoregressive inference. The 2023 GQA paper reported quality close to conventional multi-head attention with speed comparable to multi-query attention in its experiments. (arxiv.org)

Use beyond text

Multi-head attention also operates on nonlinguistic sequences. In a Vision Transformer, images are divided into patches and represented as a sequence processed by Transformer blocks. Attention relates patch representations, allowing the architecture to perform image classification without requiring convolutional feature extraction. The original Vision Transformer study demonstrated this approach through large-scale pretraining and evaluation on image-recognition benchmarks. (arxiv.org)