Ziyue (Alvin) Liu
NeurIPS 2024

CoMERA: Computing- and Memory-Efficient Training via Rank-Adaptive Tensor Optimization

Zi Yang1, Ziyue Liu2, Samridhi Choudhary3, Xinfeng Xie4, Cao Gao4, Siegfried Kunzmann3, Zheng Zhang2

1University at Albany, SUNY2University of California, Santa Barbara3Amazon Alexa AI4Meta

In short

Tensor compression can shrink a model’s weights by large factors, but its ranks are hard to choose and its many small tensor operations tend to make training slower on GPUs. CoMERA trains the model in compressed form and learns the ranks during training, balancing loss against model size. With three GPU optimizations, a six-encoder transformer trains 2–3× faster per epoch than with standard training while becoming 43× to 80× smaller, and pre-training CodeBERT-large runs about 2× faster with a 4.23× smaller model.

Time per training epoch (min)
six-encoder transformer on MNLI, one RTX 3090
Standard trainingCoMERA
05101520
14.7
7.2
14.7
6.4
15.3
5.5
3264128
batch size
Show the numbers
BatchStandardStandardwith CUDA GraphCoMERAwithout CUDA GraphCoMERA
3214.718.514.37.2
6414.716.68.26.4
12815.316.47.45.5
Standard training is shown without CUDA Graph, its faster setting here.
Model size (MB) and validation accuracy
the same transformer on MNLI
Standard training
256 MB62.2%
CoMERA, early stage
5.9 MB63.3%
Late stage, target 0.8
4.9 MB62.2%
Late stage, target 0.5
3.9 MB62.1%
Late stage, target 0.2
3.2 MB61.5%
0100200
Whole-model size. The late stage trains toward a target ratio for the tensorized layers.
2–3×
faster per training epoch than standard training
43×
smaller transformer after early-stage training, with higher accuracy
99×
smaller DLRM recommendation model, with 7× less peak memory
9×
less memory than GaLore in single-batch training

Abstract

Training large AI models such as LLMs and DLRMs costs massive GPUs and computing time. The high training cost has become only affordable to big tech companies, meanwhile also causing increasing concerns about the environmental impact. This paper presents CoMERA, a Computing- and Memory-Efficient training method via Rank-Adaptive tensor optimization. CoMERA achieves end-to-end rank-adaptive tensor-compressed training via a multi-objective optimization formulation, and improves the training to provide both a high compression ratio and excellent accuracy in the training process. Our optimized numerical computation (e.g., optimized tensorized embedding and tensor-vector contractions) and GPU implementation eliminate part of the run-time overhead in the tensorized training on GPU. This leads to, for the first time, 2–3× speedup per training epoch compared with standard training. CoMERA also outperforms the recent GaLore in terms of both memory and computing efficiency. Specifically, CoMERA is 2× faster per training epoch and 9× more memory-efficient than GaLore on a tested six-encoder transformer with single-batch training. Our method also shows ∼2× speedup than standard pre-training on a BERT-like code-generation LLM while achieving 4.23× compression ratio in pre-training. With further HPC optimization, CoMERA may reduce the pre-training cost of many other LLMs.

Learn the ranks while training

CoMERA keeps every large weight compressed from the first step. A linear layer’s weight is reshaped into a high-order tensor and stored as a tensor train: a chain of small cores, linked by dimensions called ranks that set both the size and the capacity of the layer. Embedding tables, whose shapes are very unbalanced, use the related tensor-train-matrix format.

Fixed ranks have to be guessed in advance. CoMERA instead places a diagonal matrix D on every link between cores and treats training as a problem with two goals, low loss L and a small model S. The rank of a link is the number of nonzero entries on its diagonal.

G1G2G3n1n2n3D1rank 3D2rank 2
nonzero entryzero entry, slice dropped
A tensor-train layer with rank-control diagonals. Each link’s rank is the number of nonzero entries on its diagonal. A zero entry drops the matching slice of both neighboring cores, and a layer whose ranks all reach zero is removed.

In the early stage, an ℓ1 penalty on the diagonals, a convex stand-in for counting nonzeros, pushes entries to zero while the model learns. An ℓ2 penalty on the cores keeps them from growing to make up for shrinking diagonals, and with it the early-stage solution is Pareto-optimal: among models whose cores stay within a size bound, none has both a lower loss and a smaller relaxed size (Proposition 3.1).

early:  min  L(G, D) + γ Ŝ(D) + β ‖G‖2late:  min  max{w1(L − L0), w2(S − S0)} + ρ(L + Ŝ) + β ‖G‖2

An optional late stage steers the model toward a chosen loss L0 and size S0: each step works on whichever of the two is further from its target. It can also be run on a model that is already trained.

Making small tensor operations fast on GPUs

A compressed model needs far fewer operations, but they come as many small tensor contractions that GPUs handle poorly. Three changes turn the saving into real speed.

  1. Embedding lookup without repeated work. A batch looks up many repeated rows, and different rows share parts of their tensor index. CoMERA keeps only unique rows, contracts neighboring cores in pairs so each shared part is computed once, and assembles the rows with one batched matrix multiply: 4–5× faster and 2–3× less memory than a plain tensor-train-matrix lookup. Section 4.1
  2. One contraction order for both passes. CoMERA multiplies the input-side cores together and the output-side cores together, contracts the input through them, and keeps these intermediate results for the backward pass. For large batches this order is close to the fewest operations possible (Proposition 4.1). Section 4.2
  3. CUDA Graph. Capturing the whole training step as one CUDA Graph removes most of the cost of launching many small kernels. At batch 32, an epoch drops from 14.3 to 7.2 minutes. Section 4.3

Results

Most runs use one NVIDIA RTX 3090. The main model is a six-encoder transformer trained on MNLI with sequence length 128; its embedding table and all linear layers are tensorized, starting from rank 30.

43× to 80× smaller at about the same accuracy

Early-stage training shrinks the tensorized layers 74× and the whole model 43×, and it reaches 63.3% validation accuracy against 62.2% for standard training. The late stage reaches a chosen size with a small loss: at the smallest target the tensorized layers are 361× smaller and the model 80× smaller, at 61.5%. Some layers lose all their ranks and are removed; at that target, the whole second-to-last encoder and some linear layers in other encoders are gone.

Six-encoder transformer on MNLI
Validation accuracyWhole model (MB)Tensorized layers (MB)
Standard training62.2%2561×2531×
CoMERA, early stage63.3%5.943×3.474×
CoMERA, late stage, target 0.862.2%4.952×2.4105×
CoMERA, late stage, target 0.562.1%3.965×1.4181×
CoMERA, late stage, target 0.261.5%3.280×0.7361×

Target ratios are for the tensorized layers.

The more it compresses, the faster it trains

Time per epoch falls as the compression ratio grows, and CoMERA is faster than standard training at every ratio above 1. With CUDA Graph for both, the 50× model trains 2.6× faster than standard training at batch 32. Standard training itself is faster without CUDA Graph, which forces every batch to the same sequence length, and CoMERA is still 2.0× to 2.8× faster than that.

Time per epoch (min) by compression ratio
batch 32, all runs with CUDA Graph
50× smaller
7.2
4.9× smaller
7.79
2.2× smaller
11.16
1.5× smaller
13.94
1.1× smaller
17.46
Standard training
18.5
05101520
Show the numbers
CompressionBatch 32Batch 64Batch 128
50×7.26.45.5
4.9×7.797.796.21
2.2×11.169.79.07
1.5×13.9412.4411.56
1.1×17.4615.5314.35
Standard training18.516.616.4

Faster and smaller than GaLore and LTE

With every method run under CUDA Graph, CoMERA is about 2× faster per epoch than GaLore and 3× faster than LTE, and it uses the least memory at every batch size. With a single batch, it needs 184 MB against 1,674 MB for GaLore. CoMERA and GaLore both reach about 64% validation accuracy; LTE does not converge on this task with its default settings. The same ordering holds on an RTX 4090.

Time per epoch (min)
six-encoder transformer, RTX 3090
GaLoreLTECoMERA
020406080
45.3
7.2
6.4
5.5
13264128
batch size
Show the numbers
BatchGaLoreLTECoMERA
178.3n/a45.3
3214.819.87.2
6413.816.56.4
12813.215.55.5
LTE needs a batch size that is a multiple of its 16 heads, so it has no batch-1 run.
Peak memory (MB)
six-encoder transformer, RTX 3090
GaLoreLTECoMERA
04k8k12k
184
2,174
3,800
6,998
13264128
batch size
Show the numbers
BatchGaLoreLTECoMERA
11,674n/a184
323,3584,3202,174
644,5326,4003,800
1287,71010,9226,998
LTE needs a batch size that is a multiple of its 16 heads, so it has no batch-1 run.

A 99× smaller recommendation model

On Meta’s DLRM with the Criteo Ad Kaggle data, the ten largest embedding tables are compressed in tensor-train-matrix format and the wider fully connected layers in tensor-train format. After two epochs, CoMERA matches standard training’s accuracy with a 99× smaller model and 7× less peak memory. Per epoch it is about 2× slower than standard training here, because DLRM spends its time on embedding lookups rather than matrix multiplies; the optimized lookup still cuts CoMERA’s epoch from 1,344 s to 807 s.

DLRM, batch 10,000Standard trainingCoMERA
Test accuracy78.68%78.76%
Normalized cross-entropy0.7930.792
Model size4.081 GB0.041 GB99× smaller
Peak memory18.275 GB2.612 GB7× less
Accuracy and normalized cross-entropy on the test set; peak memory includes data and back-end overhead.

Pre-training language models, a first look

For CodeBERT-large (357M) on CodeSearchNet, the CoMERA model has 84M parameters, 4.23× fewer, and pre-trains 2.3× and 1.9× faster than standard training in the two phases, with a final loss of 0.40 against 0.28. A 125M CoMERA version of BERT-large (336M) trained on Wikipedia ends at a loss of 1.45 against 1.26, and after fine-tuning it is ahead on SST-2 and MRPC and behind on SQuAD.

BERT-large after fine-tuningStandard, 336MCoMERA, 125M
SST-2, accuracy91.74%92.10%
MRPC, accuracy86.00%86.82%
SQuAD, F190.68%88.76%

Limitations

Citation

BibTeX
@inproceedings{NEURIPS2024_8d35d225,
 author = {Yang, Zi and Liu, Ziyue and Choudhary, Samridhi and Xie, Xinfeng and Gao, Cao and Kunzmann, Siegfried and Zhang, Zheng},
 booktitle = {Advances in Neural Information Processing Systems},
 doi = {10.52202/079017-2456},
 editor = {A. Globerson and L. Mackey and D. Belgrave and A. Fan and U. Paquet and J. Tomczak and C. Zhang},
 pages = {77200--77225},
 publisher = {Curran Associates, Inc.},
 title = {CoMERA: Computing- and Memory-Efficient Training via Rank-Adaptive Tensor Optimization},
 url = {https://proceedings.neurips.cc/paper_files/paper/2024/file/8d35d225230a9d77b29c1dd300e48ad9-Paper-Conference.pdf},
 volume = {37},
 year = {2024}
}