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 , , and be input matrices containing query, key, and value representations. For head , learned projection matrices form
These projections are linear maps, although implementations may also include bias terms. Each head computes scaled dot-product attention:
Here is the query/key dimension per head, denotes a matrix transpose, and 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,
Concatenation preserves the separate head outputs; the learned output projection mixes them into the layer’s output representation. (arxiv.org)
Dimensions and implementation
With query positions and key positions, each head’s attention weights have shape . Its output has 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 evenly among heads, giving head dimension . 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 positions, explicitly storing every head’s attention matrix requires elements per example. From the matrix dimensions, the score and value-aggregation operations require arithmetic when total head width is fixed; input and output projections add approximately . 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)