skip to content

Why do graph attention networks concatenate multiple heads in hidden layers but average them at the output?

level: middleimportance: should knowfreq 42%

answer

  1. heads are independent copies of the layer
  2. hidden width is a resource, output width is not
  3. the last layer's size is pinned by the task
  4. average before the final nonlinearity

basics

~20 s

Hidden layers concatenate heads because the extra width carries several independent weightings of the same neighbourhood and steadies training. The final layer must emit exactly one value per target, so its heads are averaged first and the output nonlinearity applied after.

solid answer

~50 s

Each head has its own transform and its own scoring vector, so each produces its own coefficient distribution over the same neighbourhood and its own `F'`-dimensional output. In a hidden layer you concatenate them: 8 heads of 16 features give a 128-wide node vector, several independent views of the neighbourhood survive into the next layer, and averaging the noisy single-head coefficient estimates is deferred. At the prediction layer that no longer works — the output width is pinned by the task, one logit per class, and concatenating 8 copies of a class-sized vector would produce the wrong shape. So the standard convention averages the heads elementwise first and then applies the final nonlinearity, which keeps the width right and also reduces the variance the individual heads' coefficients introduce. The cost of concatenation is that the next layer's input width, and therefore its parameter count, scales with the number of heads.

go deeper

for a junior

Remember that each head is an independent copy producing its own weighting of the neighbourhood, and that many heads are used together rather than one. Knowing which end of the network concatenates and which averages is the recall being tested.

for a middle

Explain the shape argument out loud: hidden width is free to grow, output width is pinned by the number of targets, so the heads must be collapsed at the end and averaging does it without extra parameters.

for a senior

Bring the cost side — per-edge coefficient storage multiplied by head count, and the next layer's parameters multiplied by it too — and describe how you would check whether the heads are actually learning different things before paying for more of them.

for a principal

Own head count as a capacity-versus-memory dial on a graph whose degree distribution you know, and be ready to justify trading heads for per-head width when hub nodes dominate the memory bill.

## What a head is One attention head is a complete copy of the layer's machinery: its own linear transform `W_k`, its own scoring vector `a_k`, and therefore its own softmax over each node's neighbourhood. `K` heads run on the same graph and the same inputs, and produce `K` different coefficient distributions and `K` different output vectors per node. Nothing couples them during the forward pass; they interact only through how their outputs are combined and through the gradient that flows back. ## Hidden layers — concatenate ``` h_i' = concat over k=1..K of sigma(sum_j alpha_ij^k W_k z_j) ``` With `K = 8` heads of `F' = 16` features each, a node leaves the layer as a 128-dimensional vector. Two things are gained. **Several weightings survive.** One head may concentrate on the single highest-scoring neighbour, another may spread its mass over a broader set. Concatenation keeps both intact for the next layer to use, so the network is not forced to pick one story about the neighbourhood at this depth. **Training is steadier.** A single head's coefficients are a high-variance quantity early in training — the softmax is sharp, the scorer is barely trained, and one unlucky score can dominate a node's update. With several heads, no single coefficient distribution controls the whole representation, and the empirical effect is a less brittle optimisation. The price is width. Concatenating `K` heads multiplies the next layer's input dimension by `K`, so the next layer's transform matrix grows by the same factor, and so does the activation memory carried into it. Going from 1 head to 8 does not cost 8 times the current layer — it costs 8 times the *next* layer's parameters too. ## The output layer — average ``` h_i' = sigma_out((1/K) * sum over k=1..K of sum_j alpha_ij^k W_k z_j) ``` At the prediction layer the output width is not a free hyperparameter: a 7-class node-classification task needs exactly 7 logits per node. Concatenating 8 heads that each emit 7 numbers would give 56, which is not a class distribution. Three options exist in principle — shrink each head to `classes/K` features, add a projection after concatenation, or average — and the standard convention is to average, because it is parameter-free and keeps every head predicting in the same output space where the average is meaningful. Note the ordering: the heads are averaged **before** the final nonlinearity, not after. Averaging `K` softmax distributions and taking a softmax of the averaged logits are different operations, and the convention is the latter. Averaging is also a variance reduction. Each head's prediction carries noise from its own randomly initialised scorer; a mean over `K` heads damps the part of that noise that is uncorrelated across heads. That is a real but modest benefit — the heads share inputs and gradients, so their errors are far from independent. ## Do the heads actually learn different things? Nothing enforces it. The heads differ only through random initialisation, and it is entirely possible for several of them to converge onto near-identical coefficient patterns, at which point the extra width is buying redundancy rather than diversity. This is worth checking directly: compare the coefficient distributions across heads on the same nodes, or look at the pairwise correlation of their output vectors. If they are near-duplicates, fewer heads with a wider `F'` each is usually the better trade — same width, fewer softmaxes, fewer per-edge coefficients to store. ## Memory, concretely Per layer, per head, the model stores one coefficient per edge. Eight heads on a graph with ten million edges is eighty million coefficients per layer during the backward pass. On a graph with heavy-tailed degrees this is dominated by a handful of hub nodes. Head count is therefore not a free knob: it multiplies both the width flowing forward and the per-edge state held for the gradient. ## The answer an interviewer is listening for The crisp version is that concatenation and averaging are answering two different constraints. In the middle of the network, width is a resource and you want to keep the heads distinct. At the end of the network, width is fixed by the task, so the heads must be collapsed, and averaging is the parameter-free way to collapse them while staying in the output space. A candidate who says only "averaging is more stable" has half the answer; the shape constraint at the output layer is the part that makes the convention non-negotiable.

  • What happens to the next layer's parameter count when you go from one head to eight concatenated heads?
    Its input width multiplies by eight, so its transform matrix has eight times the rows and eight times the parameters, and the activations feeding it grow the same way. The head count is not just a cost in the layer that owns the heads; it propagates into the layer downstream, which is why people usually shrink the per-head feature count when raising the head count.
  • Is anything forcing the heads to learn different neighbourhood weightings?
    No. They differ only through random initialisation and are free to converge on near-identical coefficient patterns. Check it by comparing the coefficient distributions of different heads on the same nodes, or the correlation between head outputs. If they are duplicates, the width is redundancy, and one head with a wider feature dimension gives you the same capacity for fewer softmaxes and less per-edge state.
  • Could you concatenate at the output layer instead if you wanted to?
    Only by fixing the shape yourself: give each head a fraction of the target width, or add a projection after concatenation. Both work but add either a constraint or parameters, and they break the property that every head predicts in the same output space. Averaging is the convention because it is parameter-free and shape-correct without either fix.

saying these in an interview costs you the question

  • Thinks all heads share one transform and one scoring vector
  • Says heads are averaged in hidden layers and concatenated at the output
  • Ignores that head count multiplies the next layer's parameters
  • Assumes heads are guaranteed to learn diverse patterns
  • Applies the output nonlinearity per head before averaging

context