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 () throughout the model’s depth. In this work, we investigate whether this uniform allocation limits the representational capacity available to different stages of a Transformer. To address 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 establishes a local-to-global representational hierarchy: early layers leverage fewer, wider heads to capture complex, local compositional patterns, while deeper layers deploy more, narrower heads to refine these representations into specialized linguistic features. Crucially, this structural shift is parameter-neutral and theoretically compute-neutral, preserving identical weight matrices and theoretical FLOP budgets as the standard Transformer. Across three model scales (124M, 354M, and 757M), the Prism Transformer consistently achieves lower validation loss than uniform baselines while maintaining identical measured training throughput on our 8× H100 GPU setup, alongside improved performance on downstream zero-shot benchmarks including PIQA, HellaSwag, ARC-Easy, and WinoGrande. Our findings suggest that non-uniform subspace allocation can improve the use of model capacity within the standard Transformer parameter and FLOP budget. Our code is available at https://anonymous.4open.science/r/PrismTransformer-593A/.
est. 32% chance this paper gets accepted at ICLR 2027.
What do you think this paper will get?
All positions stay anonymous.