Learning Functional Subspaces for Neural Network Compression
Abstract
Modern transformers couple impressive capabilities with substantial memory and compute demands. Low-rank weight factorization reduces both while keeping the matrices dense, and therefore efficient on standard hardware. However, several existing methods choose the subspace removed from each weight matrix by _local_ closed-form criteria—activation energy, layer-wise reconstruction error, or a quadratic loss approximation—and thus ignore how errors propagate through the rest of the network. At aggressive compression ratios, these errors compound through depth and performance collapses. The number of directions removed from each layer is set before training according to that layer's measured effect on the output. We introduce _Learnable Subspace Projections_ (LSP), which instead learns the removed subspaces end-to-end. Each linear layer, or _tied group_ of layers that read the same activations, receives an orthogonal projector. All projectors are optimized jointly against a _global_ objective, the KL divergence to the dense model's output distribution, while the pretrained weights stay frozen. Projectors are initialized from a whitened SVD truncation, and ranks are allocated by measured cost: the output KL divergence each projector induces, per parameter saved. After training, the projectors merge into standard low-rank factors, and each tied group shares one input factor, so attention can cache a single narrow latent instead of full keys and values. Across decoder-only LLMs (OPT-125M/1.3B, Qwen3-4B, Llama-2-7B) and a Vision Transformer (ViT-B/16), LSP outperforms activation-based, curvature-based, and rank-learning baselines by a margin that grows with the compression ratio. At 70% compression, LSP brings Llama-2-7B to WikiText-2 perplexity versus for the strongest baseline, and to mean zero-shot accuracy versus . The factorized model decodes up to faster than the dense model at small batch sizes, and caching the shared latent reduces the combined memory of weights and KV cache by at a 128k-token context, versus at most for untied factorizations.
est. 32% chance this paper gets accepted at ICLR 2027.
What do you think this paper will get?
All positions stay anonymous.