Multi-Head Attention Residuals
AuthorsCheng Luo, Zefan Cai, Junjie Hu
Resources
MHAR lets different feature groups retrieve different layers of a transformer's history, improving language-model quality with little added cost.
Key results
MHAR's reduction versus the standard Transformer at 350M parameters.
H=8 is the large-scale routing-head setting adopted by the authors.
Relative training throughput versus the baseline Transformer.
Accuracy-point improvement for the 8B MHAR model over schedule-matched continued pretraining.
Approximate token scale used for the 8B mid-training experiments.
What the paper found
Cheng Luo, Zefan Cai, and Junjie Hu introduce Multi-Head Attention Residuals, or MHAR, a modification of attention residuals that lets each feature subspace retrieve different layers from a Transformer’s depth history. Instead of one learned query producing a single softmax over prior sublayer outputs, MHAR reshapes that query into independent heads, creating block-diagonal depth routing with zero additional parameters and negligible FLOPs; H=1 exactly reproduces the single-head method. Trained from scratch with a Qwen3-style architecture on NVIDIA’s deduplicated, quality-filtered anneal_pt_v3 corpus, MHAR reduces validation loss versus a standard Transformer by 0.061, 0.149, and 0.140 at 100M, 350M, and 1B parameters, respectively, outperforming hyper-connections and single-head attention residuals. The best routing granularity is not unlimited: validation loss is U-shaped, with H=4 or H=8 consistently optimal, while H=16 over-splits coherent feature groups. Fused Triton kernels raise 1B-model training throughput to 0.55 times the baseline while keeping near-baseline memory. An identity-preserving delta conversion also scales the method to Marin-8B, a Llama-3.1-8B model, during mid-training on 1.9T tokens, adding 3.2 GSM8K points and 3.1 GPQA points over schedule-matched continued pretraining. The authors’ query probes support their central explanation: as models widen, feature subspaces increasingly disagree about which layers to read, making a shared routing distribution a measurable bottleneck.
Original abstract
Transformers propagate information across depth through a single additive residual stream: every sublayer reads only the most recent state. Attention residuals relax this by letting each sublayer attend, through a learned softmax. However, that read uses a single query shared across the entire width, so every feature subspace must read the depth history through one distribution. The cost of this forced compromise grows with how much the subspaces disagree about which layers to read, and disagreement grows with model width. We introduce Multi-Head Attention Residuals (MHAR): the routing query is reshaped into H per-subspace heads, each with its own softmax over the depth history. The read becomes block-diagonal, the reshape adds zero parameters and negligible compute, and H = 1 recovers attention residuals exactly. Trained from scratch on a deduplicated Nemotron-based anneal corpus that is quality-filtered and STEM- and code-heavy, MHAR improves validation loss over a standard Transformer at 100M, 350M, and 1B (-0.061, -0.149, and -0.140). It achieves the best result among four methods in every setting, with the gain increasing from 100M to the larger scales. The head count is a real design axis rather than a free knob: validation loss is U-shaped with respect to H, with a flat optimum at H = 4 or H = 8 across scales. We adopt H = 8 for large-scale models; over-splitting beyond this point (H = 16) consistently gives back part of the gain. A direct probe of the trained queries confirms that learned subspace disagreement is the underlying driver. Fused Triton routing kernels increase attention-residual training throughput from 0.2-0.5x to 0.55-0.88x of the baseline while maintaining near-baseline peak memory. An identity-preserving conversion using delta attention residuals supports 8B mid-training, yielding improvements of +3.2 on GSM8K and +3.1 on GPQA.
Read the original paperMore in Transformers
Browse all 42 papers →Pretraining Latent Information Feedback Transformers with Teacher Supervision
Dor Tirosh, Ido Amos, Mor Geva
LIFT teaches Transformers to pass rich hidden-state information across steps, potentially making language models more efficient and capable than standard feed-forward designs.
The Geometry of Inference in Transformer Residual Streams
Timur Mudarisov, Mikhail Burtsev, Radu State
This paper shows how Transformer hidden states gradually geometrically converge toward the correct prediction while eliminating competing possible outcomes.
Transformers Stop Thinking Too Early, and a Tiny LoRA Fixes It
Zehao Jin, Ruixuan Deng, Junran Wang
A small LoRA update appears to make transformers carry information through many more layers, dramatically extending their ability to follow long chains without retraining the full model.