Prism Transformer: Progressive Head Schedules for Hierarchical Attention Processing
Abstract
Multi-head attention conventionally partitions the hidden dimension equally across all heads at every layer, enforcing an identical representational subspace dimension (dh = dmodel/h) throughout the model's depth. In this work, we identify this uniform allocation as a fundamental structural bottleneck: due to their restricted dimensional space, early-layer heads are unable to faithfully capture complex, high-dimensional contextual patterns. To resolve this, we introduce the Prism Transformer, a novel architectural paradigm that replaces the static, uniform head configuration with a progressive head schedule. By monotonically increasing the head count across layers, the Prism Transformer naturally establishes a local-to-global representational hierarchy: early layers leverage fewer, exceptionally wide heads to capture complex, long-range contextual patterns, while deep layers deploy many, narrow heads to decompose these patterns into specialized linguistic features. Crucially, this structural shift is parameter-neutral, compute-neutral, and introduces zero training or inference overhead, preserving identical weight matrices and FLOP budgets as the standard Transformer. Across three model scales (124M, 354M, and 757M), the Prism Transformer consistently outperforms uniform baselines, achieving significant reductions in validation perplexity alongside uniform gains on downstream zero-shot benchmarks (including HellaSwag, PIQA, WinoGrande, BLiMP, ARC-Easy, and WikiText). Our findings demonstrate that non-uniform subspace allocation unlocks latent capacity within the standard Transformer budget, offering a more effective use of model capacity within the standard Transformer budget.