LaX: Boosting Low-Rank Training of Foundation Models via Latent Crossing
University of California, Santa Barbara*Equal contribution
Low-rank models save compute but lose quality, because every layer is limited to a small rank. LaX adds a residual connection between the low-rank latents of consecutive layers, so each layer can use more than one low-rank subspace at almost no cost. With LaX, low-rank models match or beat full-rank training with 2–3× fewer parameters, and LoRA fine-tuning improves as well.
Show the numbers
| Model | Params (M) | Accuracy (%) | + LaX params (M) | + LaX accuracy (%) |
|---|---|---|---|---|
| ViT-B, full-rank | 86.56 | 76.74 | ||
| Tensor train | 41.18 | 71.11 | 41.44 | 75.43+4.32 |
| SVD | 44.17 | 75.20 | 44.24 | 77.20+2.00 |
| CoLA | 44.17 | 76.04 | 44.24 | 77.84+1.80 |
| ViT-L, full-rank | 304.33 | 77.10 | ||
| Tensor train | 101.97 | 75.21 | 102.10 | 77.77+2.56 |
| SVD | 115.77 | 76.81 | 115.92 | 78.60+1.79 |
| CoLA | 115.77 | 77.63 | 115.92 | 79.07+1.44 |
Abstract
Training foundation models such as ViTs and LLMs requires tremendous computing cost. Low-rank matrix or tensor factorization offers a parameter-efficient alternative, but often downgrades performance due to the restricted parameter space. In this work, we introduce Latent Crossing (LaX) – a simple yet effective plug-and-play module that enhances the capacity of low-rank models by enabling information flow across low-rank subspaces. We extensively validate the benefits of LaX on pre-training tasks with ViT-Base/Large and LLaMA-like models ranging from 60M to 1B parameters. LaX boosts low-rank model performance to match or exceed the full-rank baselines while using 2-3× fewer parameters. When equipped with low-rank adapters (i.e., LoRA) for fine-tuning LLaMA-7/13B, LaX consistently improves performance on arithmetic and common sense reasoning tasks with negligible cost.
A residual path between latent spaces
A low-rank layer replaces a weight W with factors: a down-projection Ai to r dimensions and an up-projection Bi back out. This saves parameters and compute, but everything the layer passes on has to fit through the rank-r latent, and raising r gives the savings back. LaX keeps r and adds the latent that the same kind of layer in the previous block has already computed:
Written as one matrix, the layer now acts on its own input and the previous layer's input, and its output can use a second low-rank subspace. Since hi−1 is produced by the forward pass anyway, this plain version adds no parameters. In a transformer, LaX links attention to attention and MLP to MLP across blocks. A LayerNorm on the LaX path keeps training stable; without it the loss becomes NaN right away.
SVD / CoLA
Between the rank-r latents of the same layer type in consecutive blocks. CoLA adds a nonlinearity σ there.
Tensor train
Across layers, and inside one layer between results of the core-contraction chain.
LoRA
Between the adapters of consecutive layers. The frozen weights are untouched.
Gates
When the two latents have different sizes, or to weigh the residual, a small gate G sits on the path. The overhead stays small because r is small; only the Dense Gate grows with r2.
| Gate | Extra parameters | Extra FLOPs per block | Where the paper uses it |
|---|---|---|---|
| Identity | none | O(1) | LLM pre-training, commonsense fine-tuning |
| Linear | one scalar per path | O(nr) | arithmetic fine-tuning |
| Tensor | two small cores | O(nr) | ViT pre-training |
| Dense | an r × r matrix | O(nr2) | ablations |
On SVD ViT-B every gate helps at every rank, and the Tensor Gate helps most: 2.00, 1.88 and 1.68 points at ranks 256, 128 and 64, for at most 0.07M extra parameters.
| SVD ViT-B | rank 256 | rank 128 | rank 64 |
|---|---|---|---|
| Without LaX | 75.2044.17M | 74.4722.94M | 71.2612.32M |
| Identity | 75.81+0.61 | 74.65+0.18 | 71.64+0.38 |
| Linear | 76.11+0.91 | 75.31+0.84 | 71.72+0.46 |
| Tensor | 77.20+2.00 | 76.35+1.88 | 72.94+1.68 |
| Dense | 77.03+1.83 | 75.43+0.96 | 72.33+1.07 |
Results
Vision transformers
ViT-B and ViT-L trained from scratch on ImageNet-1K for 300 epochs, with the Tensor Gate. LaX lifts all three low-rank models. On ViT-B, SVD and CoLA with LaX beat the full-rank model with about half its parameters, and tensor train gains the most, 4.32 points. On ViT-L all three beat the full-rank model with 2.6 to 3.0× fewer parameters. The accuracy chart at the top of this page shows these runs.
Language model pre-training
LLaMA-style models from 60M to 1B parameters trained on C4 with compute-optimal token budgets. LaX lowers perplexity for both SVD and CoLA at every size. Plain SVD falls behind most baselines, and with LaX it comes close to or beats LORO and CoLA. CoLA with LaX beats full-rank training at all four sizes: at 1B it reaches 14.78 against 15.56 with 609M parameters instead of 1,339M, using the Identity Gate. LaX also steadies training, so both SVD and CoLA can use learning rates about ten times larger.
Show all methods, with parameter counts
| Method | LLaMA-60M | LLaMA-130M | LLaMA-350M | LLaMA-1B |
|---|---|---|---|---|
| Full-rank | 34.0658M | 24.36134M | 18.80368M | 15.561,339M |
| ReLoRA | 37.0458M | 29.37134M | 29.08368M | 18.331,339M |
| GaLore | 34.8858M | 25.36134M | 18.95368M | 15.641,339M |
| SLTrain | 34.1544M | 26.0497M | 19.42194M | 16.14646M |
| LORO | 33.9643M | 24.5994M | 18.84185M | 15.19609M |
| SVD | 36.2543M | 26.8494M | 21.18185M | 16.54609M |
| SVD + LaX | 33.5444M | 24.6394M | 18.90185M | 15.51609M |
| CoLA | 34.0443M | 24.4894M | 19.40185M | 15.52609M |
| CoLA + LaX | 33.2144M | 24.2199M | 18.51196M | 14.78609M |
LoRA fine-tuning
LLaMA-7B and 13B fine-tuned with rank-32 LoRA on arithmetic and commonsense reasoning, following the LLM-Adapters setup, with LaX between the adapters of consecutive layers. The share of trainable parameters stays at 0.83% and 0.67%. The average rises on both benchmarks and both models; the largest single gain is 6.6 points on HellaSwag for 7B, and the only drop is BoolQ on 13B, by 0.8.
Show every task
| Task | 7B LoRA | 7B + LaX | 13B LoRA | 13B + LaX |
|---|---|---|---|---|
| MultiArith | 95.0 | 96.9+1.9 | 95.2 | 97.3+2.1 |
| GSM8K | 36.1 | 37.7+1.6 | 47.5 | 49.0+1.5 |
| AddSub | 84.3 | 84.8+0.5 | 86.0 | 86.3+0.3 |
| AQuA | 17.7 | 19.3+1.6 | 18.2 | 20.9+2.7 |
| SingleEq | 84.4 | 87.8+3.4 | 89.8 | 91.9+2.1 |
| SVAMP | 51.8 | 53.6+1.8 | 54.6 | 58.3+3.7 |
| Average | 61.6 | 63.4+1.8 | 65.2 | 67.3+2.1 |
| Task | 7B LoRA | 7B + LaX | 13B LoRA | 13B + LaX |
|---|---|---|---|---|
| BoolQ | 68.9 | 69.6+0.7 | 72.1 | 71.3−0.8 |
| PIQA | 80.7 | 81.9+1.2 | 83.5 | 85.4+1.9 |
| SIQA | 77.4 | 78.9+1.5 | 80.5 | 81.3+0.8 |
| HellaSwag | 78.1 | 84.7+6.6 | 90.5 | 91.3+0.8 |
| WinoGrande | 78.8 | 80.8+2.0 | 83.7 | 84.1+0.4 |
| ARC-e | 77.8 | 79.8+2.0 | 82.8 | 84.4+1.6 |
| ARC-c | 61.3 | 64.8+3.5 | 68.3 | 71.8+3.5 |
| OBQA | 74.8 | 78.0+3.2 | 82.4 | 83.1+0.7 |
| Average | 74.7 | 77.3+2.6 | 80.5 | 81.6+1.1 |
Citation
@inproceedings{zhang2025lax,
title = {{LaX}: Boosting Low-Rank Training of Foundation Models via Latent Crossing},
author = {Zhang, Ruijie and Liu, Ziyue and Wang, Zhengyang and Zhang, Zheng},
booktitle = {Advances in Neural Information Processing Systems},
year = {2025}
}