Learning Functional Subspaces for Neural Network Compression
Learnable Subspace Projections compress LLMs end-to-end, beating baselines at 70% compression on Llama-2-7B.
Learnable Subspace Projections (LSP) compress transformers by learning which weight subspaces to discard, instead of using local closed-form criteria that ignore how errors propagate with depth. Orthogonal projectors are optimized jointly against output KL or the original loss while pretrained weights stay frozen, then merged into standard low-rank factors. Across OPT-125M/1.3B, Qwen3-4B, Llama-2-7B, and ViT-B/16, LSP's advantage grows with compression. At 70% compression, Llama-2-7B reaches 10.9 WikiText-2 perplexity and 42.2% mean zero-shot accuracy versus 13.3 and 36.0% for the strongest baseline, with up to 1.6x faster decoding and 13.5x lower combined weight and KV-cache memory at a 128k context.
- LSP learns orthogonal projectors end-to-end while pretrained weights stay frozen.
- At 70% compression, Llama-2-7B reaches 10.9 perplexity and 42.2% zero-shot accuracy.
- Caching a shared latent shrinks weights plus KV cache 13.5x at 128k tokens.
- Factorized models decode up to 1.6x faster than the dense model at small batches.
Full article276 words · extracted from huggingface.co · click to collapse
Modern transformers pair impressive capabilities with substantial memory and compute demands. Low-rank weight factorization reduces both while keeping the matrices dense, and thus efficient on standard hardware. Existing methods, however, choose the subspace to remove from each weight matrix with local closed-form criteria: activation energy, layer-wise reconstruction error, or a quadratic approximation of the loss. These criteria ignore how errors propagate through the network, so at high compression the errors compound with depth and performance collapses. We introduce Learnable Subspace Projections (LSP), which instead learns the subspaces to discard end-to-end. Each linear layer, or tied group of layers that read the same activations, is assigned an orthogonal projector. All projectors are optimized jointly against a global objective--the KL divergence to the dense model's output distribution or the model's original training loss--while the pretrained weights remain frozen. Projectors are initialized from a whitened SVD truncation, and ranks are allocated by the output KL each projector induces per parameter saved. After training, the projectors merge into standard low-rank factors, with each tied group sharing one factor. In attention, this also lets the model cache one narrow latent in place of full keys and values. Across LLMs (OPT-125M/1.3B, Qwen3-4B, Llama-2-7B) and ViT-B/16, LSP outperforms baselines, and its advantage widens as compression increases. At -70% compression, LSP brings Llama-2-7B to 10.9 WikiText-2 perplexity and 42.2% mean zero-shot accuracy, versus 13.3 and 36.0% for the strongest baseline. The factorized model decodes up to 1.6x faster than the dense model at small batch sizes, and aching the shared latent shrinks the combined memory of weights and KV cache by 13.5x at a 128k-token context, versus at most 6.5x for untied baseline factorizations.
Text extracted automatically; images, tables and formatting may be missing. Original: https://huggingface.co/papers/2609.40127