🤖 AI Summary
This study addresses the substantial computational overhead incurred by repeatedly calculating attention weights during Transformer training and the difficulty of reusing static attention patterns. To mitigate this, we propose Selective Attention Freezing (SAF), which identifies low-variance attention heads and replaces them with fixed patterns, enabling the dynamic introduction of static attention mid-training for the first time. This approach integrates absolute and relative positional preference modeling, linear memory optimization, and fused CUDA kernels to accelerate computation. Consequently, SAF reduces memory complexity from quadratic to linear. Empirical evaluations on 124M and 1B parameter models demonstrate speedups of 1.056× and 1.068×, respectively, with only a marginal 0.77% increase in perplexity. By significantly accelerating long-sequence fine-tuning and prefilling phases, this work effectively balances computational efficiency with model performance.
📝 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.