about
Your Transformer is Secretly Linear (arxiv.org)
37 points by skilled on May 23, 2024 | hide | past | pdf | 6 comments on HN

In plain words: A language model's layers pass along signals linked by a near-perfect straight-line mapping, since each layer adds a small change to what it receives. Deleting or straight-line-replacing the most linear layers barely hurt performance, and a penalty making layers less linear improved small-model scores.

Abstract

This paper reveals a novel linear characteristic exclusive to transformer decoders, including models such as GPT, LLaMA, OPT, BLOOM and others. We analyze embedding transformations between sequential layers, uncovering a near-perfect linear relationship (Procrustes similarity score of 0.99). However, linearity decreases when the residual component is removed due to a consistently low output norm of the transformer layer. Our experiments show that removing or linearly approximating some of the most linear blocks of transformers does not affect significantly the loss or model performance. Moreover, in our pretraining experiments on smaller models we introduce a cosine-similarity-based regularization, aimed at reducing layer linearity. This regularization improves performance metrics on benchmarks like Tiny Stories and SuperGLUE and as well successfully decreases the linearity of the models. This study challenges the existing understanding of transformer architectures, suggesting that their operation may be more linear than previously assumed.

Anton Razzhigaev, Matvey Mikhalchuk, Elizaveta Goncharova, Nikolai Gerasimenko, Ivan Oseledets, Denis Dimitrov, Andrey Kuznetsov
arXiv:2405.12250 · cs.LG, cs.AI, cs.CL · submitted May 19, 2024
abstract · pdf · html · 9 pages, 9 figures

add comment on HN

The authors hypothesize at the bottom of page 3 that linear layers can combine to form nonlinear functions. This is wrong, but maybe I’m misunderstanding what they are trying to say.
I think they're talking about linearity as transformations with a linearity score close to 1. They defined that linearity score a little higher up. Such that composing many almost linear transformations will create a total transformation that is very nonlinear.
A fun example of that (making a neural network using floating point error as a source of non-linearity): https://youtu.be/Ae9EKCyI1xU?si=n9vgvCvxoxrQeKd8
That is not 100% what I read in this paper. There are several takes:

1. LoRA makes transformers linear versus pre-training that keeps non-linearity (in 3.1 and 3.2). What is kinda to be expected.[One more insight is that the combination of seemingly linear blocks can lead to non-linear output]. Thus you can replace part of the layers in fine-tuned models by nn.Linear for inference ...

2. There is a way to make LoRA keep non-linearity by changing loss function and improve performance of the model (in 4, Cosine Similarity regularization term).

3. Small models are surprisingly unaffected. But IMO that may be because of small number of layers and weights overall and the adapter layer being much larger in comparison to the model size.

Scaling linear algebra in the end is probably all we’ll need in the end. Only missing data and compute to get there
Memory capacity is a much bigger problem. Mixtral 8x22B is a 200GB+ model.