Ziyue (Alvin) Liu
NeurIPS 2025

LaX: Boosting Low-Rank Training of Foundation Models via Latent Crossing

Ruijie Zhang*, Ziyue Liu*, Zhengyang Wang, Zheng Zhang

University of California, Santa Barbara*Equal contribution

In short

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.

ImageNet-1K top-1 accuracy (%), trained from scratch
hollow: low-rank model · filled: with LaX · dashed: full-rank model
ViT-Bfull-rank: 86.56M parameters, 76.74%
Tensor train
full-rank71.1175.43
SVD
75.2077.20
CoLA
76.0477.84
707274767880
ViT-Lfull-rank: 304.33M parameters, 77.10%
Tensor train
full-rank75.2177.77
SVD
76.8178.60
CoLA
77.6379.07
707274767880
Show the numbers
ModelParams (M)Accuracy (%)+ LaX params (M)+ LaX accuracy (%)
ViT-B, full-rank86.5676.74
Tensor train41.1871.1141.4475.43+4.32
SVD44.1775.2044.2477.20+2.00
CoLA44.1776.0444.2477.84+1.80
ViT-L, full-rank304.3377.10
Tensor train101.9775.21102.1077.77+2.56
SVD115.7776.81115.9278.60+1.79
CoLA115.7777.63115.9279.07+1.44
+4.32
accuracy points for tensor-train ViT-B on ImageNet-1K, from 71.11% to 75.43%
2.2×
fewer parameters than full-rank LLaMA-1B, at lower perplexity: 14.78 against 15.56
+2.6
points of average commonsense accuracy for LoRA on LLaMA-7B, with the same trainable parameters

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:

hi = Ai xi,   hi−1 = Ai−1 xi−1ỹi = Bi (hi + hi−1) = [ BiAi   BiAi−1 ] [ xi ; xi−1 ]

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

outputsoutputsinputsAirBiAi−1rBi−1GLaX

Between the rank-r latents of the same layer type in consecutive blocks. CoLA adds a nonlinearity σ there.

Each panel shows layer i−1 below layer i. Blue arrows are LaX paths; G is the optional gate.

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.

GateExtra parametersExtra FLOPs per blockWhere the paper uses it
IdentitynoneO(1)LLM pre-training, commonsense fine-tuning
Linearone scalar per pathO(nr)arithmetic fine-tuning
Tensortwo small coresO(nr)ViT pre-training
Densean r × r matrixO(nr2)ablations
n is the sequence length and r the rank; the low-rank block itself costs O(ndr + n2d).

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-Brank 256rank 128rank 64
Without LaX75.2044.17M74.4722.94M71.2612.32M
Identity75.81+0.6174.65+0.1871.64+0.38
Linear76.11+0.9175.31+0.8471.72+0.46
Tensor77.20+2.0076.35+1.8872.94+1.68
Dense77.03+1.8375.43+0.9672.33+1.07
ImageNet-1K top-1 accuracy (%). Small numbers show the gain over the model without LaX; under “Without LaX” they show its parameter count.

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.

Validation perplexity on C4, lower is better
hollow: without LaX · filled: with LaX · dashed: full-rank
LLaMA-60Mfull-rank 34.06
SVD
full-rank36.2533.54
CoLA
34.0433.21
3334353637
LLaMA-130Mfull-rank 24.36
SVD
full-rank26.8424.63
CoLA
24.4824.21
24252627
LLaMA-350Mfull-rank 18.80
SVD
full-rank21.1818.90
CoLA
19.4018.51
18192021
LLaMA-1Bfull-rank 15.56
SVD
full-rank16.5415.51
CoLA
15.5214.78
151617
Show all methods, with parameter counts
MethodLLaMA-60MLLaMA-130MLLaMA-350MLLaMA-1B
Full-rank34.0658M24.36134M18.80368M15.561,339M
ReLoRA37.0458M29.37134M29.08368M18.331,339M
GaLore34.8858M25.36134M18.95368M15.641,339M
SLTrain34.1544M26.0497M19.42194M16.14646M
LORO33.9643M24.5994M18.84185M15.19609M
SVD36.2543M26.8494M21.18185M16.54609M
SVD + LaX33.5444M24.6394M18.90185M15.51609M
CoLA34.0443M24.4894M19.40185M15.52609M
CoLA + LaX33.2144M24.2199M18.51196M14.78609M
Full-rank, ReLoRA, GaLore, SLTrain, LORO and CoLA numbers are from earlier papers; SVD and the LaX rows are from this paper.

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.

Average accuracy (%), LoRA with rank 32
hollow: LoRA · filled: LoRA with LaX
Arithmetic reasoning, 6 tasks
LLaMA-7B
61.663.4
LLaMA-13B
65.267.3
6062646668
Commonsense reasoning, 8 tasks
LLaMA-7B
74.777.3
LLaMA-13B
80.581.6
7476788082
Show every task
Arithmetic reasoning
Task7B LoRA7B + LaX13B LoRA13B + LaX
MultiArith95.096.9+1.995.297.3+2.1
GSM8K36.137.7+1.647.549.0+1.5
AddSub84.384.8+0.586.086.3+0.3
AQuA17.719.3+1.618.220.9+2.7
SingleEq84.487.8+3.489.891.9+2.1
SVAMP51.853.6+1.854.658.3+3.7
Average61.663.4+1.865.267.3+2.1
Commonsense reasoning
Task7B LoRA7B + LaX13B LoRA13B + LaX
BoolQ68.969.6+0.772.171.3−0.8
PIQA80.781.9+1.283.585.4+1.9
SIQA77.478.9+1.580.581.3+0.8
HellaSwag78.184.7+6.690.591.3+0.8
WinoGrande78.880.8+2.083.784.1+0.4
ARC-e77.879.8+2.082.884.4+1.6
ARC-c61.364.8+3.568.371.8+3.5
OBQA74.878.0+3.282.483.1+0.7
Average74.777.3+2.680.581.6+1.1

Citation

BibTeX
@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}
}