Static attention heads speed up training with minor accuracy tradeoff
When Can Attention Heads Be Statically Defined?
Machine LearningComputation and Language
Summary
Some parts of AI language models called attention heads learn similar patterns repeatedly, which means their work can be reused instead of recalculated every time. The authors discovered a way to pick which attention heads have stable patterns and then fix those patterns halfway through training to save computing time. This approach speeds up training and fine-tuning with only a small drop in how well the model predicts text. It also helps the model handle longer input sequences better after further adaptation.
What this means in practice
- •For machine learning engineers: Reduce training time of large language models by statically fixing stable attention heads during training.
- •For ai infrastructure teams: Accelerate model fine-tuning and long-context inference by integrating fixed-pattern attention heads to reduce computation.
Authors
Weixian Waylon Li, Yintao Tai, Marcio Fonseca, Shay B. Cohen
Abstract
Some attention heads learn similar patterns across inputs. Reusing these patterns could reduce training cost by avoiding repeated query-key score computation and softmax. Through controlled pretraining comparisons, we identify Selective Attention Freezing (SAF), which selects heads with low attention-pattern variance and replaces their attention weights with fitted post-softmax means halfway through training. We represent these fixed patterns with absolute-position and relative-distance preferences, reducing storage from quadratic to linear in sequence length. A fused kernel reconstructs the patterns and executes ordinary-attention and replaced heads together. At matched training-token budgets, replacing 25% of attention heads gives 1.056x faster post-replacement optimiser updates at 124M parameters and 4K context, with a 0.77% perplexity increase. At 1B and 8K context, post-replacement updates are 1.068x faster on four GPUs including communication, with a 0.51% perplexity increase. The resulting models also accelerate long-input finetuning and causal prefill. After associative-recall adaptation, the 124M model with 25% replacement generalises to more key-value pairs at a fixed length better than ordinary attention and two pruning controls.