Learning Functional Subspaces for Neural Network Compression
Authors: Massimo Bini, Anders Christensen, Stephan Alaniz, Judah Goldfeder, Ole Winther, Yann LeCun, Ravid Shwartz-Ziv, Zeynep Akata
Organizations: Helmholtz Munich · Technical University of Munich · MCML · Orbital Industries · LTCI, T´el´ecom Paris, Institut Polytechnique de Paris · Columbia University · University of Copenhagen · Technical University of Denmark · New York University
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.
Figures & tables
Figure 1: Spectral energy is not function. Each arrow is a direction in the input activation space of a linear layer W : its length is the spectral energy it carries, its orientation how much the network output depends on it (horizontal right: most), and faded arrows are removed. Activation-based low-rank compression (top) keeps the highest-energy directions, even when the output barely depends on them (red), and can discard functional ones (blue). LSP (bottom) learns the projector P against the network output, removing the directions the output is least sensitive to, whatever their energy.
Figure 2: Overview of LSP. (1) An orthogonal projector P=I−UU⊤ removes the learned subspace span(U) from the input or output of W . (2) Layers that read the same activations (Q/K/V, gate/up) share a tied projector; the remaining (untied) layers are projected on their smaller side. (3) Calibration activations give the whitened SVD of WS , whose trailing directions initialize U . (4) The KL divergence that each candidate truncation induces on the model output, per saved parameter, sets how many directions each unit removes. (5) The unconstrained V is trained, with W frozen, on the KL to the dense model or on the task loss; U=qf(V) is its orthonormal QR factor, so the removed directions are those the output depends on least. (6) The trained projectors decompose into low-rank factors U⊥U⊥⊤ , and are merged with the frozen weights.
Comp.
Method
OPT-125M
OPT-1.3B
Qwen3-4B
Llama2-7B
Dense (Base)
27.9
14.7
8.0
5.5
-30%
activations/ [-2.5pt]reconstruc.
NoLSP
48.5
23.0
16.8
9.4
ASVD
106.8
396.1
154.5
98.2
SliceGPT
51.7
23.9
15.0
10.3
SVD-LLM (W)
48.7
20.2
19.1
10.4
Swift-SVD
47.0
19.8
15.6
10.4
Table 1: Wikitext-2 perplexity ( ↓ ) of LLMs at −30/50/70% . One of the two LSP variants has the lowest perplexity in all 12 settings, by a margin that grows with the ratio; at −70% both variants beat every baseline on all four models.
Llama-2-7B (Alpaca calibration)
Qwen3-4B (C4 calibration)
Method
Comp.
ARC-e
PIQA
Openb.
WinoG.
HellaS.
MathQA
Avg.
Comp.
ARC-e
PIQA
Openb.
WinoG.
HellaS.
MathQA
Avg.
Dense (Base)
–
75.6
77.8
32.8
69.9
57.1
28.1
56.9
–
79.0
77.9
32.0
70.3
54.6
53.8
61.3
SVD-LLM (W)
-30%
61.4
67.2
24.4
61.3
38.7
23.0
46.0
-20%
59.0
70.5
22.4
59.4
40.1
26.5
46.3
Swift-SVD
61.6
68.1
24.0
61.5
38.9
23.0
46.2
63.6
72.2
26.8
64.2
45.7
28.6
50.2
Dobi-SVD
53.3
64.1
20.4
57.5
35.7
22.8
42.3
57.5
69.0
22.6
62.9
41.5
26.7
46.7
SVD-LLM
68.5
72.7
30.2
64.7
49.0
24.3
51.6
63.1
72.7
27.2
64.9
47.6
31.2
51.1
Table 2: Zero-shot accuracy ( ↑ , %). Left: Llama-2-7B at 30/50/70% compression, every method calibrated on Alpaca. Right: Qwen3-4B at 20/40/60% compression, C4 calibration. Both LSP variants lead every baseline in mean accuracy at every ratio, by a margin that grows with the ratio.
Single pool
Diverse pool
Comp.
Method
Source
Transfer
Source
Transfer
Gain
Dense (Base)
89.9
47.4
89.9
47.4
—
-30%
SliceGPT
86.3
38.5
85.1
40.0
+1.5
SVD-LLM (W)
88.2
41.2
87.7
42.3
+1.1
FLAR-SVD
88.8
43.6
88.6
44.7
+1.1
PELA
89.2
42.3
88.9
43.9
+1.6
Table 3: ViT-B/16 compressed on a single pool ( ∼ 47k CIFAR-100 images) or a diverse pool ( ∼ 47k images from CIFAR-100, Food-101, CIFAR-10, EuroSAT, STL-10 and DTD). Source : CIFAR-100 accuracy; Transfer : mean linear-probe accuracy on Pets, Aircraft and Places365, in neither pool; Gain : diverse minus single transfer, computed before rounding.
Table 6
Figure 3: Rank allocation and spectral removal on Llama-2-7B. Left: fraction of full rank retained by each unit after compression, by block and projection type; measured-KL allocation is far from uniform. Middle: amplitude Δi lost along each singular direction of the original weights; average over layers; Right: difference to NoLSP, σi(aiLSP−aiNoLSP) ; the leading direction changes most, and outside Q/K/V LSP removes more of it and retains more of the trailing directions, most visibly in the down projection. All panels show runs at −30/50/70% compression (light to dark). Crosses give the delta for σ1 . Plots for more models and layers in Appendix D .
Appendix figures & tables15 assets
Supplementary material from the paper’s appendix.
Appendix
Algorithm 1 LSP.
Model
Metric
Method
-30%
-50%
-70%
OPT-125M
ppl ↓
LSP T
24.48 ± 0.11
30.63 ± 0.28
42.41 ± 0.06
LSP
30.95 ± 0.42
35.35 ± 0.67
48.88 ± 0.24
OPT-1.3B
ppl ↓
LSP T
12.93 ± 0.02
16.42 ± 0.11
23.33 ± 0.26
LSP
15.70 ± 0.15
18.51 ± 0.11
22.39 ± 0.30
Qwen3-4B
ppl ↓
LSP T
9.20 ± 0.03
12.48 ± 0.02
19.40 ± 0.14
LSP
9.13 ± 0.07
11.42 ± 0.10
16.18 ± 0.08
Appendix
Table 6: Seed sensitivity of LSP and LSP T : mean ± standard deviation over three full-pipeline seeds, each run read at its own validation-selected epoch. The LLMs are WikiText-2 perplexity, ViT-B/16 is CIFAR-100 accuracy on the single (CIFAR-100 only) pool.
−30%
−50%
−70%
Model
Dense
Whitened
Plain
Whitened
Plain
Whitened
Plain
OPT-125M
27.6
582
282
664
534
1,538
2,211
OPT-1.3B
14.6
26.1
52.5
74.4
571
814
18,786
Qwen3-4B
7.9
40.1
29.4
255
125
7,186
4,449
Llama-2-7B
5.5
9.2
13.8
22.1
36.9
172
215
Appendix
Table 7: Initialization before training. WikiText-2 test perplexity of the training-free truncation at uniform allocation, from the whitened truncation of Section 2.2 ( Whitened , NoLSP at uniform allocation) or from the trailing directions of the plain activation Gram ( Plain ), removing the same number of parameters per unit.
−30%
−50%
−70%
OPT-125M
Whitened
Plain
Whitened
Plain
Whitened
Plain
Untrained
582
282
664
534
1,538
2,211
Trained (LSP)
31.7
32.3
35.9
36.6
44.8
45.8
Appendix
Table 8: Initialization before and after training, OPT-125M. WikiText-2 test perplexity (dense 27.6 ) of the training-free truncation and of LSP distilled from it, only the initialization changed.
−30%
−50%
−70%
Initialization
Q/K/V, gate/up
PPL
Avg6
PPL
Avg6
PPL
Avg6
Whitened
tied
12.2
52.5
21.8
47.3
37.9
42.0
Whitened
untied
14.8
50.4
21.3
47.0
38.2
41.5
Plain
tied
16.4
50.6
30.7
45.3
49.8
39.3
Plain
untied
22.2
47.5
27.9
44.0
47.6
38.8
Appendix
Table 9: Initialization and tying after training, Llama-2-7B. LSP distilled on Alpaca, uniform removed rank per layer, the better of two learning rates per arm; WikiText-2 test perplexity and mean zero-shot accuracy (Avg6) of the validation-selected epoch.
−30%
−50%
−70%
Allocation
PPL
Avg6
PPL
Avg6
PPL
Avg6
Uniform retained parameters ( Equation 19 )
16.5
41.1
18.7
37.8
24.2
35.0
Uniform removed rank
16.8
40.6
18.9
37.6
23.4
35.4
Diagonal Fisher
16.7
41.4
18.6
38.1
25.6
34.5
Measured KL ( ours )
15.6
42.4
17.9
38.8
23.2
36.3
Appendix
Table 10: Rank allocation, OPT-1.3B. Four budgets at the same realized savings (within 0.05 points), whitened tied initialization, distilled on WikiText-2, the better of two learning rates per arm; WikiText-2 test perplexity and Avg6 of the validation-selected epoch.
Allocation
Comp.
ARC-e
PIQA
Openb.
WinoG.
HellaS.
MathQA
Avg.
Llama-2-7B (Alpaca calibration)
Dense (Base)
–
75.6
77.8
32.8
69.9
57.1
28.1
56.9
Uniform
-30%
70.3
74.5
30.6
63.4
50.0
25.5
52.4
Measured KL
71.2
75.4
30.0
64.0
50.4
26.4
52.9
Uniform
-50%
64.3
70.6
25.8
59.8
43.8
24.8
48.2
Measured KL
65.2
70.8
24.6
58.6
43.4
23.9
47.8
Appendix
Table 11: Uniform against measured-KL allocation, zero-shot accuracy ( ↑ , %). Uniform is the retained fraction of Equation 19 ; Measured KL is the per-unit KL cost of Equation 3 . Same whitened initialization, distillation objective (LSP) and realized budget; the Qwen3-4B pair also differs in K/V routing ( Table 12 ). Llama-2-7B is calibrated on Alpaca and compressed by −30/50/70% , Qwen3-4B on C4 by −20/40/60% , as in Table 2 .
−20%
−40%
−60%
Allocation
Tied group
Avg6
PPL
KV
Avg6
PPL
KV
Avg6
PPL
KV
Uniform
Q/K/V, input
53.4
15.9
1357
49.4
18.4
1018
43.3
23.0
679
Uniform
K/V, output
54.0
16.0
1242
49.3
18.6
932
42.9
23.6
622
Measured KL
Q/K/V, input
56.4
15.0
1822
49.3
17.9
1701
42.9
23.2
1282
Measured KL
K/V, output
56.9
14.7
2048
50.9
17.1
2048
44.2
22.4
1920
Appendix
Table 12: Allocation and K/V routing, Qwen3-4B. LSP distilled on C4, same learning rate, epoch budget and realized savings; gate/up are tied in every arm. Q/K/V, input : one shared input projection for Q, K and V. K/V, output : the grouped-query routing, with K and V sharing an output-side factor and Q projected alone. The uniform Q/K/V row is that table’s Uniform cell. Avg6 and C4 test perplexity of the validation-selected epoch; KV is the floats cached per token per layer (dense 2048 ). Best per column in bold.
Figure 4: Rank allocation across models. Retained rank r as a fraction of the full rank d , by block and projection type, for the measured-KL LSP checkpoints of OPT-125M, OPT-1.3B, Qwen3-4B and the CIFAR-100 ViT-B/16; the same panel for Llama-2-7B is Figure 3 (left). Shades stack the three ratios, −30% lightest behind and −70% darkest in front, so each shade’s top edge is that ratio’s profile wherever the allocations nest; r/d=1 is a unit the allocator left dense. The removal grid is 7 points on OPT-125M, OPT-1.3B and the ViT and 15 on Qwen3-4B. Qwen3-4B has grouped-query attention, so Q is projected alone and K/V share an output-side factor, and each gets its own panel.
Figure 5: Spectral removal in single blocks. The two spectral panels of Figure 3 for an early, a middle and a late block (columns) of OPT-1.3B (blocks 3 , 12 , 21 of 24 , top) and Llama-2-7B (blocks 4 , 16 , 28 of 32 , bottom), on the measured-KL checkpoints of Table 1 . For each model, the first row is the amplitude lost along each original singular direction, Δi=σi(ai−1) , with −σi of the uncompressed layer dashed, and the second the difference to NoLSP, σi(aiLSP−aiNoLSP) , at the same measured-KL ranks. Each row shares one vertical axis. Shades are the −30/50/70% runs (light to dark), the bands are a running mean over 9 directions, a missing band is a unit left dense at that ratio, and tied Q/K/V and up/gate panels average their member matrices; crosses give the raw value on direction 1, parked on the axis edge when outside it.
Memory
Compute and speed (batch 1)
Method
Weights
KV floats
Weights+KV
Max ctx
Prefill
Decode
Decode tok/s
(GiB)
/token/layer
@128k, b8 (GiB)
(k tok, b8)
TFLOPs @2k
GFLOP/step
@512
@4k
@16k
Llama-2-7B (dense)
12.55
8192
524.6
19.7
29.26
14.29
142
85
36
−30%
SVD-LLM (W)
8.93
2866 ( 2.9× )
188.1 ( 2.8× )
59.0
21.30
10.40
135
83
35
Swift-SVD
8.93
2866 ( 2.9× )
188.1 ( 2.8× )
59.0
21.30
10.40
137
83
36
Appendix
Table 13: Inference efficiency of compressed Llama-2-7B at every ratio, the checkpoints of Table 1 , on one GH200 (bf16, SDPA). Memory , analytic from the checkpoints’ ranks (factors vs. dense): weights, floats cached per token per layer with the latent KV cache ( z=Ax ; a unit the measured-KL allocation leaves uncompressed stays dense and caches full-width K/V), weights plus KV cache at 128k tokens and batch 8, and the longest context that fits 95.5 GiB at batch 8. Speed : FLOPs to prefill 2k tokens and per decode step at 2k context, and measured CUDA-graph-compiled decode tok/s at batch 1 for 512, 4k and 16k tokens of context. Best factorized method per ratio in bold.
−30%
−50%
−70%
Allocation
KV
Max ctx
tok/s
KV
Max ctx
tok/s
KV
Max ctx
tok/s
Llama-2-7B (dense)
8192
19.7
142
8192
19.7
142
8192
19.7
141
Uniform
1673
101.0
157
1195
145.6
165
717
249.8
209
Measured KL
2316
73.0
172
1096
158.8
193
556
322.1
221
Appendix
Table 14: Deployment cost of the allocation, Llama-2-7B. The uniform and measured-KL allocations compared in Table 11 , here on the WikiText-2-calibrated checkpoints, on one GH200, Q/K/V tied on the input side. KV : floats cached per token per layer, and Max ctx : the longest context (k tokens) that fits 95.5 GiB at batch 8, both analytic from the checkpoints’ ranks, counting a unit the allocator left uncompressed as caching full-width K,V. tok/s : measured CUDA-graph-compiled decode throughput at batch 1 and 512 tokens of context, both allocations in the same session, best of two passes. Within a ratio both hold the same weights and issue the same decode FLOPs.
−20%
−40%
−60%
Allocation
Tied group
KV
Max ctx
tok/s
Avg6
KV
Max ctx
tok/s
Avg6
KV
Max ctx
tok/s
Avg6
Qwen3-4B (dense)
2048
74.6
175
61.3
2048
74.6
175
61.3
2048
74.6
175
61.3
Uniform
Q/K/V, input
1357
114.5
147
53.4
1018
155.1
170
49.4
679
236.3
186
43.3
Uniform
K/V, output
1242
125.1
144
54.0
932
169.4
165
49.3
622
257.9
179
42.9
Measured KL
Q/K/V, input
1822
85.3
180
56.4
1701
92.8
187
49.3
1282
125.2
201
42.9
Measured KL
K/V, output
2048
75.9
176
56.9
2048
77.1
182
50.9
1920
83.6
193
44.2
Appendix
Table 15: Deployment cost of allocation and K/V routing, Qwen3-4B. The four arms of Table 12 on one GH200. KV : floats cached per token per layer, and Max ctx : the longest context that fits 95.5 GiB at batch 8, both analytic from the merged checkpoints’ ranks, counting a unit the allocator left uncompressed as caching full-width K,V. tok/s : measured CUDA-graph-compiled decode throughput at batch 1 and 512 tokens of context, best of two passes. Within a ratio all four arms hold the same weights and issue the same decode FLOPs, given in the text.
Figure 6: Compression cost with measured-KL allocation on three LLMs. WikiText-2 perplexity at −30/50/70% compression (marker size) against wall-clock hours. LSP costs include initialization, training through the selected epoch, and the full KL measurement on the 7-point grid on OPT-125M and the 15-point grid on Qwen3-4B and Llama-2-7B.
Compressed on
Transfer to
Source
Method
CIFAR-100
Food-101
CIFAR-10
EuroSAT
STL-10
DTD
Images
Pets
Aircraft
Places365
Mean
CIFAR-100
Dense (Base)
—
71.9
36.9
33.5
47.4
89.9
-30% parameters
FLAR-SVD
✓
10k
67.1
32.3
31.2
43.6
88.7
✓
✓
20k
68.6
34.1
32.2
45.0
88.8
✓
✓
✓
✓
40k
68.6
33.2
32.7
44.9
88.7
Appendix
Table 16: Complete downstream-transfer results for compressed ViT-B/16, for every calibration pool and target. Compressed on : the datasets in the pool (✓) and its size, the same pools for every method; ✓ † is CIFAR-100 alone at the size of the six-dataset pool. Transfer to : linear-probe accuracy on Pets, Aircraft and Places365, absent from every pool, and their mean; Source : CIFAR-100 accuracy through the original head. LSP T needs CIFAR-100 labels. Best per column, pool and ratio in bold .
Pipeline parallelism enables training of large language models that exceed single-device memory, yet inter-stage activation communication becomes the dominant bottleneck when trained on low-bandwidth networks. Recent work in this area has proposed using fixed orthogonal projections to compress activations. However, this still results in a significant performance degradation and requires a number of non-standard adaptations to constrain the optimization. A natural alternative is to learn a low rank projection for each pipeline stage, however maintaining the necessary orthogonality of these projectors during training remains a challenge. We present Manifold Aware Projection Learning (MAPL), a method that treats inter-stage compression as a learnable orthogonal projection under explicit Stiefel manifold (orthogonal matrices) constraints. Rather than prescribing a fixed global subspace, MAPL lets each pipeline stage discover and continuously adapt its own task-optimal compression subspace via manifold-constrained steepest descent. To recover token-specific signals at stage boundaries, we introduce per-stage factorized anchor embeddings that allow for full-rank activation reconstruction with negligible communication overhead. We further show that we can incorporate residual vector quantization after projection with a streaming codebook synchronization protocol that amortizes dictionary communication. Across LLaMA models from 150M to 1B parameters we show that MAPL can be easily applied to the existing pipeline and can achieve high compression with neglibile performance degradation with a drastically improved tradeoffs in performance vs. compression compared to Subspace Networks.
Paul Janson, Edouard Oyallon, Eugene Belilovsky
Concordia University · Mila Quebec AI Institute · CNRS, Sorbonne University
Recent SVD based compression methods for large language models like SVD LLM and Basis Sharing can be unified under one optimization problem. While mathematical proofs and tests on Pythia models show this unified approach improves weight reconstruction error by up to 46% percent it fails in practical tasks. Downstream metrics like perplexity and accuracy severely degrade compared to standard per layer SVD LLM. The authors explain this failure mechanistically. Although the bundle method mathematically couples adjacent layers the transformer residual stream actually decouples them during forward passes. Thus per layer optimality matters more than joint cross layer optimization. The paper concludes that weight space reconstruction is a flawed objective for cross layer compression and future methods must focus on per layer activation reconstruction instead.
Low-rank decomposition is a promising compression paradigm for large language models (LLMs), yet its effectiveness hinges on rank budget allocation across weight matrices: uniform or hand-crafted rules ignore module-wise importance, while learning-based allocation incurs substantial training overhead. We formulate rank allocation as a global sorting-and-truncation pipeline that scores every singular component by combining local singular energy with global functional importance, estimated via layer-wise input--output cosine similarity on a tiny calibration set. We show, both geometrically and empirically, that high input--output cosine similarity implies low effective rank. We further propose rank-preserving fine-tuning (RPFT), which adapts only a small subset of retained singular components so that the allocated rank stays bounded without re-decomposition. Experimental results show that UniRank cuts zero-shot perplexity by up to 50%, improves average reasoning accuracy by 3.0% over LoRAP at 25% sparsity, and boosts four SVD-based decomposition methods as a plug-and-play module.
Chao Han, Yongjie Du, Junjie Tan +1
Ningbo Institute of Digital Twin, Eastern Institute of Technology, Ningbo