Attention is Not All You Need: Pure Attention Loses Rank Doubly Exponentially with Depth

Yihe Dong, Jean-Baptiste Cordonnier, Andreas Loukas

Introduction

The attention mechanism [BCB15] was initially developed to better learn long-range sequential knowledge, and found effective use in transformer networks [VSP+17]. Since then, attention-based architectures have permeated across data domains machine learning applications, such as in natural language processing [DCLT18, POT16], speech recognition [LZL+20], and computer vision [RPV+19, BZV+19]. As such, it is vital to develop tools to understand the inner workings of transformers and attention in general, both to shed light on existing models, and to design more effective future models.

This work provides new insights about the operation and inductive bias of networks built by stacking multiple self-attention layers. Surprisingly, we find that pure self-attention networks (SANs), i.e., transformers with skip connections and multi-layer perceptrons (MLPs) disabled, lose expressive power doubly exponentially with respect to network depth. More specifically, we prove that the output converges with a cubic rate to a rank one matrix that has identical rows. While we derive the convergence bounds in part by using properties of stochastic matrices, our results go beyond what one would expect based on standard results. In particular, by leveraging the cascading effects of specifically stacking self-attention modules, we show exponentially faster convergence than what standard theory prescribes. Furthermore, while previous studies have considered the rank of individial self-attention matrices [WLK+20, KVPF20, CLJ20a], our results are the first to address conditions under which the entire network converges to rank one.

This raises the question, why do transformers work? Our analysis indicates that skip connections play a key role in mitigating rank collapse, and MLPs can slow down the convergence by increasing their Lipschitz constant. We characterize these counteracting forces by proving upper and lower bounds of this convergence behavior under SAN architectural variants that resemble transformers. Our results reveal a previously unknown vital utility of skip connections, beyond facilitating optimization and gradient flow [HZRS16a, BFL+18].

In the process, we develop a new path decomposition to study self-attention networks. Namely, we decompose a SAN into a linear combination of weakly-interdependent paths, where each ‘path’ corresponds to a deep single-head SAN. Intuitively, one can view the self-attention heads in each layer of the original network as different gateways, and a path follows a sequence of gateway choices, one gateway per layer (Figure 1). Coupled with the rank collapse analysis, our results suggest that deep SANs with skip connections behave like an ensemble of weakly-dependent shallow networks.

Our main contributions are as follows: (1) We present a systematic study of building blocks of the transformer, revealing opposing impacts between self-attention and the counteracting forces: skip connections and MLP, in contributing and preventing a rank collapse in transformers. As a corollary, this reveals a previously unknown vital effect of skip connections beyond facilitating optimization. (2) We propose a new method for analyzing SANs via a path decomposition, revealing SANs as an ensemble of shallow networks. (3) We verify our theory with experiments on common transformer architectures.

Attention doubly exponentially loses rank

We start by studying self-attention networks (SANs) built exclusively out of multi-head self-attention layers. We prove that SANs converge exponentially (with depth) to a rank-1 matrix that makes all tokens identical.

Our analysis in §2.1 relies on an unconventional way to express the output of a multi-head SAN as a sum of single-head networks. We refer to the latter as paths, where each path is denoted by a sequence of attention heads (see Figure 1). A proof sketch of why rank collapse occurs is given in §2.2, whereas the main rank collapse result is presented in §2.3.

Let X\bm{X} be a n×dinn\times d_{\textit{in}} input matrix consisting of nn tokens. A SAN is built out of LL multi-head self-attention layers, each having HH heads. The output of the hh-th self-attention head can be written as

Above, WV,h\bm{W}_{V,h} is a din×dvd_{\textit{in}}\times d_{\textit{v}} value weight matrix and the n×nn\times n row-stochastic matrix Ph\bm{P}_{h} is given by

where (1) the key and query weight matrices WK,h\bm{W}_{K,h} and WQ,h\bm{W}_{Q,h} are of size din×dqkd_{\textit{in}}\times d_{\textit{qk}}, (2) WQK,h=WQ,hWK,h⊤\bm{W}_{QK,h}=\bm{W}_{Q,h}\bm{W}_{K,h}^{\top}, and (3) the softmax operates independently on each row of its input. We obtain the final equation by noting that softmax is shift-invariant and disregarding terms that provide a constant contribution across rows [CLJ20a].

The output of each SAN layer is formed by concatenating the individual outputs of all HH attention heads (along the last dimension) and linearly projecting them onto a subspace of appropriate size:

where we set Wh=WV,h WO,h⊤\bm{W}_{h}=\bm{W}_{V,h}\,\bm{W}_{O,h}^{\top} and bO=∑hbO,h\bm{b}_{O}=\sum_{h}\bm{b}_{O,h}.

Let Xl\bm{X}^{l} be the output of the ll-th layer and fix X0=X\bm{X}^{0}=\bm{X}. As is common practice, we let all layers consist of the same number of heads.

Excluding biases 1bO,h⊤\text{{1}}\bm{b}_{O,h}^{\top}, the SAN output is given by

which, after unrolling the recursion backwards, yields:

The above equations have a clear interpretation if we think of the SAN as a directed acyclic graph, with nodes corresponding to self-attention heads and directed edge connecting heads of consecutive layers.

We formalize this intuition in the following:

The output of a depth LL self-attention network with HH heads per layer (including biases and skip connections) is given by

where Ppath=PhLL⋯Ph11\bm{P}_{\textit{path}}=\bm{P}_{h_{L}}^{L}\cdots\bm{P}_{h_{1}}^{1} is an input-dependent stochastic matrix, whereas Wpath=Wh11⋯WhLL\bm{W}_{\textit{path}}=\bm{W}_{h_{1}}^{1}\cdots\bm{W}_{h_{L}}^{L} and b\bm{b} do not depend on the input.

The proof follows from the fact that the set of row-stochastic matrices is closed under multiplication (i.e., PhLL⋯Phii\bm{P}_{h_{L}}^{L}\cdots\bm{P}_{h_{i}}^{i} is row-stochastic) and, moreover, for any row-stochastic matrix P\bm{P}, we have P1=1\bm{P}\text{{1}}=\text{{1}}. ∎

Each term in (1) describes a path of length LL across heads of different layers

and there are a total of HLH^{L} such paths without skip connections.

The path decomposition thus describes the action of a multi-head SAN as the combination of simpler single-head networks. To gain intuition on path interdependence, it helps to split the operations performed into two types: those that act across tokens (multiplication from left) and those that apply independently on each token (multiplication from right). As seen, though paths can interact through token mixing (since PpathP_{\textit{path}} matrices jointly depend on XX), token-wise operations are independent. We can also notice that biases are not particularly meaningful: their total contribution amounts to the single term 1b⊤\text{{1}}\bm{b}^{\top} independently of the number of layers or heads used.

In the following we show that each path converges rapidly (as a function of length) to a rank-1 matrix with identical rows. Interestingly, this convergence is so dominant that adding more layers to the SAN does not help: though the number of paths is increased exponentially, each path degenerates doubly exponentially, leading also to a rank-1 output.

2 Convergence of single-head SAN

Before tackling the full SAN, it is instructive to consider the behavior of each path separately. We examine, in particular, how the residual

As the following result shows, the residual norm converges to zero surprisingly quickly (doubly exponentially with a cubic rate):

For any single-head SAN consisting of LL layers with ∥WQKl∥1∥WVl∥1,∞≤β\|\bm{W}_{QK}^{l}\|_{1}\|\bm{W}_{V}^{l}\|_{1,\infty}\leq\beta and for a term γ\gamma that depends on the attention entries, we have that

which amounts to a doubly exponential convergence to a rank-1 matrix.

For the full theorem, we refer the reader to the Appendix.

Note that the bound in Eq 2 guarantees ∥res(SAN(X))∥1,∞\|\text{res}(\text{SAN}(\bm{X}))\|_{1,\infty} convergence for all inputs of small residual whenever 4γβ<dqk4\gamma\beta<\sqrt{d_{\textit{qk}}}. In practice, our experiments imply that the region for convergence can be much greater.

The identified cubic rate of convergence is significantly faster than what would be expected when analyzing products of stochastic matrices (linear rate). As a rule of thumb, to achieve a decline of three orders of magnitude, say from 1000 to 1, one could expect a linear rate of convergence to require roughly a dozen iterations, whereas a cubic rate can do so in just two or three iterations. The reason why we get a cubic rate is that the rank of attention matrices depends also on the rank of the input. As we show, the self-attention heads mix tokens faster when formed from a low-rank matrix. This phenomenon becomes stronger as we build deeper SANs, leading to a cascading effect.

We provide a proof sketch bellow. Detailed proofs can be found in the Appendix.

To analyze how the formation of Ph\bm{P}_{h} is affected by the rank of the input, we start by writing X=1x⊤+R\bm{X}=\text{{1}}\bm{x}^{\top}+\bm{R} for R=res(X)\bm{R}=\text{res}(\bm{X}) and expanding the attention matrix accordingly:

Invoking once more the shift-invariance property of the softmax, the above can be simplified to

for some appropriate r\bm{r}. Observe that if the matrix within the softmax was 1r⊤\text{{1}}\bm{r}^{\top}, then Ph\bm{P}_{h} would also degenerate to a rank-1 matrix: softmax(1r⊤)=1 q⊤\text{softmax}(\text{{1}}\bm{r}^{\top})=\text{{1}}\,\bm{q}^{\top} and the convergence would happen instantly.

The proof builds on this observation by showing that if E=RWQKdqkR⊤\bm{E}=\bm{R}\frac{\bm{W}_{QK}}{\sqrt{d_{\textit{qk}}}}\bm{R}^{\top} is small then Ph\bm{P}_{h} is almost rank-1:

where D\bm{D} is diagonal and Dii=max⁡j∣δi⊤E(δj−δj′)∣\bm{D}_{ii}=\max_{j}|\bm{\delta}_{i}^{\top}\bm{E}(\bm{\delta}_{j}-\bm{\delta}_{j^{\prime}})|. Thus, we have

and, moreover, ∥res(PhX)∥≤2 ∥D1 q⊤R∥.\|\text{res}(\bm{P}_{h}\bm{X})\|\leq 2\,\|\bm{D}\text{{1}}\,\bm{q}^{\top}\bm{R}\|. The proof concludes by bounding the above term and applying the argument recursively over successive layers. ∎

3 Exponential convergence for attention networks

We now move on to analyse the convergence of SANs with multiple heads per layer.

Consider a depth-LL and width-HH self-attention network without skip connections. Suppose that ∥WQK,hl∥1∥Whl∥1,∞≤β\|\bm{W}_{QK,h}^{l}\|_{1}\|\bm{W}_{h}^{l}\|_{1,\infty}\leq\beta for all heads h∈[H]h\in[H] and layers l∈[L]l\in[L], and let γ\gamma be a term that depends on the attention entries. We have

which amounts to a doubly exponential rate of convergence.

The bound guarantees convergence of SAN(X)\text{SAN}(\bm{X}) to rank one when 4γβH<dqk4\gamma\beta H<\sqrt{d_{\textit{qk}}}. Our experiments show that this is a rather pessimistic estimate, as, in practice, we observe widespread convergence of output to rank-1.

Remark 1. Implications for Xformers. There has been a surge of architectural variants – that we collectively refer to as Xformers – aimed to improve the vanilla transformer [VSP+17] by reducing the quadratic self-attention complexity. The rank collapse result of Theorem 2.3 carries interesting implications for these architectures. One such variant relies on low-rank or kernel-based approximations to the full attention matrix [KVPF20, WLK+20, CLD+20], in which case the paths likely converge even faster to rank one due to the imposed low-rankedness. Another variant only computes a subset of the attention matrix entries using particular patterns [ZGD+20, CGRS19], such as random patterns, in which case one expects the paths to converge more slowly, as randomization tends to increase the rank of the output.

Mechanisms that counteract rank collapse

Our findings raise a pertinent question—why do attention-based networks work in practice if attention degenerates to a rank-1 matrix doubly exponentially with depth? Aiming to obtain a deeper understanding, we focus on the transformer architecture [VSP+17] and expand our analysis by incorporating the three important components of transformers that SANs lack: skip connections, multi-layer perceptrons, and layer normalization.

We adopt a methodical approach where the modifications to the SAN architecture are introduced one at a time. For each case, we re-derive the convergence bounds and discuss the observed effect.

A simple modification to the path decomposition argument for SAN suffices to take into account skip connections. Specifically, we indicate the event that a path has skipped a layer by setting h=0h=0 on the corresponding notation:

where we have fixed P0=I\bm{P}_{0}=\bm{I} and W0=I\bm{W}_{0}=\bm{I}.

As observed, skip connections dramatically diversify the path distribution. Denote by Pl\mathcal{P}_{l} the set of paths of length ll. With skip connections enabled, we have

paths of length ll (whereas before we had only length LL paths). We hypothesize that it is the presence of short paths that stops SAN from degenerating to rank-1.

While we can derive an upper bound for the residual similar to above (which we do in the Appendix for completeness) such an upper bound is vacuously large. Indeed, it is more informative to have a lower bound on the residual, to align with practice, where SANs with skip connections do not suffer rank collapse. We present the following simple lower bound:

Consider a depth-LL and width-HH self-attention network with skip connections. There exist infinitely many parameterizations for which ∥res(XL)∥≥∥res(X)∥\|\text{res}(\bm{X}^{L})\|\geq\|\text{res}(\bm{X})\|. The preceeding holds even for L→∞L\to\infty and β\beta arbitrarily small.

The proof is elementary: by the path decomposition, there is always a path that skips all layers, i.e. the path with length 0, preserving the residual. It then follows that, for any parametrization that renders the contribution of the SAN layers orthogonal to the input, we will have ∥res(XL)∥≥∥res(X)∥\|\text{res}(\bm{X}^{L})\|\geq\|\text{res}(\bm{X})\|. A simple example of such a parametrization can be recovered by setting WVl=0\bm{W}_{V}^{l}=0 for every l∈[L]l\in[L], in which case ∥res(XL)∥=∥res(X)∥\|\text{res}(\bm{X}^{L})\|=\|\text{res}(\bm{X})\|.

A tight lower bound to the residual in the presence of skip connections is highly nontrivial, and we pose it as an open challenge to the community.

Remark 2. SANs as ensembles of shallow networks. It can be deduced from Theorem 2.3 that the SANs with skip connections enabled heavily rely on short paths (since the residual rapidly declines as the path length becomes larger). In other words, SANs behave like ensembles of shallow single-head self-attention networks. The phenomenon was previously identified for ResNets [VWB16a] (though the latter study didn’t study the rank-collapse phenomenon). Here, the components of this ensemble are inter-dependent, as each attention head participates in many paths of different lengths. Experimental results in §4 support this implication. The supplementary material also provides a study of the paths distribution across several common architectures.

2 Multi-layer perceptrons (MLP) help

We now study how using an MLP affects the residual. In particular, we focus on SANs with layers written as

Note that, to keep the notation compact, we use flf_{l} to denote both the MLP as well as the output bias.

Consider a depth-LL and width-HH SAN with MLP. Suppose that ∥WQK,hl∥1∥Whl∥1,∞≤β\|\bm{W}_{QK,h}^{l}\|_{1}\|\bm{W}_{h}^{l}\|_{1,\infty}\leq\beta for all h∈[H]h\in[H] and l∈[L]l\in[L], let γ\gamma be a term that depends on the attention entries, and fix λl,1,∞≤λ\lambda_{l,1,\infty}\leq\lambda. We have that

which amounts to a doubly exponential rate of convergence.

As seen, though the effect of MLP is less drastic than that of skip connections, the convergence rate in Cor 3.2 can be controlled by the Lipschitz constants λf,1,∞\lambda_{f,1,\infty} of the MLPs: the more powerful the MLPs are the slower the convergence becomes. This reveals a tug-of-war between the self-attention layers and the the MLPs, which due to their nonlinearity can increase the rank. §4 shows that indeed MLPs counteract convergence in experiments.

We should emphasize that using MLPs to counteract the rank-collapse is not without drawbacks: While increasing the Lipschitz constants slows down residual convergence, it also renders the model less robust and more sensitive to input perturbations [CKSN18]. Larger Lipschitz constants may also pose greater challenges to optimization, as they lead to larger gradient variance.

3 Layer normalization plays no role

Layer normalization is accomplished by rescaling and shifting the input across the feature dimension:

where bLN\bm{b}_{\text{LN}} is the mean of each column SA(X)\text{SA}(\bm{X}) and DLN\bm{D}_{LN} is a diagonal matrix with entries corresponding to the (possibly scaled or shifted) standard deviation of each column SA(X)\text{SA}(\bm{X}).

Experiments

Our experiments first test the rank collapse phenomenon in several well-known transformers architectures (§4.1). We also visually illustrate the inductive bias of some architectural variants of transformers with a toy example in §4.2 and test the paths effectiveness with respect to length in §4.3. Additional results can be found in the Appendix.

To verify our theoretical predictions, we examine the residual of three well-known transformer architectures: BERT [DCLT18], Albert [LCG+19], and XLNet [YDY+19]. Figure 2 plots the relative residual ∥res(SAN(Xl)∥1,∞/∥SAN(Xl)∥1,∞,{\|\text{res}(\text{SAN}(\bm{X}^{l})\|_{1,\infty}}/{\|\text{SAN}(\bm{X}^{l})\|_{1,\infty}}, of each layer’s output before and after the networks have been trained. To compute these ratios we ran the network on 32 samples of 128 tokens excerpts of biographies from Wikipedia [LGA16] and display the mean and standard deviation.

The experiments confirm that, as soon as the skip connections are removed, all networks exhibit a rapid rank collapse. Though MLPs do not seem to help in the mitigation of convergence, we caution that the observation is not an accurate portrayal of how trained transformers behave: removing the skip connections introduces a drastic distribution shift in the MLP input. We expect that the convergence will slow down if the network is retrained.

2 Visualizing the bias of different architectures

To empirically investigate the inductive bias of the different components of the transformer architecture, we study the behavior of a single-layer transformer when applied recurrently (akin to the universal transformer [DGV+19]) to predict a simple 2D circular sequence.

Since this recurrent application of a single-layer transformer can be reparametrized to be equivalent to a multi-layer transformer without skip connections, we hypothesize that at inference time the predicted trajectories of the two arcs will converge to the same point (indicating a rank collapse), rather than following the training trajectories. Note that the setting has also been intentionally constructed to enable training even without skip connections (by using teacher forcing) and thus to disentangle the two distinct benefits of skip connections: their ability to improve optimization and their mitigation of rank collapse.

We trained the network until it could perfectly memorize the next step on the circular trajectories with near-zero loss. Figure 3 demonstrates the trajectories predicted at inference time (i.e., without teacher forcing). As seen on the top row, without MLP or skip connections the network exhibits rank collapse. Theorem 2.2 predicts that the convergence is slower when β≥∥WQKl∥1∥WVl∥1,∞\beta\geq\|\bm{W}_{QK}^{l}\|_{1}\|\bm{W}_{V}^{l}\|_{1,\infty} increases. Indeed, as the hidden dimension increases from 32 to 128 (leading to larger β\beta at initialization), the convergence slows down, becoming hardly observable for dimension 128.

We conclude that, in accordance to our analysis, adding MLP or skip connections either stops or drastically slows down rank collapse. As observed, skip connections tend to slow down points from moving. The latter phenomenon is because in this setting skip connections introduce a bias towards remaining in the same position. On the other hand, adding MLPs does not exhibit the same bias.

3 Path effectiveness

SANs can be seen as ensembles of paths of different lengths (from 0 to LL), each involving a different sequence of self-attention heads. Our analysis of SAN with skip connections indicates that path expressivity decreases with path length, even if the number of non-linear operations involved increases. To test this hypothesis, we isolate paths of different lengths and evaluate their predictive power.

Tasks. We considered the following three tasks to test path effectiveness with respect to length:

Sequence memorization. To solve this task, a model needs to memorize a pre-determined mapping from natural language sentences and random label sequences of the same length. We use random tokens (rather than actual labels) to make this purely a test of expressiveness of a network by way of memorizing training data, rather than confounding effects such as generalizability. The models tested are trained to minimize the cross entropy loss between predicted and the ground truth labels. The training data consist of 500 English sentences from Wikipedia and News sources [DGM06, WSM+19], which are tokenized using the SentencePiece tokenizer [KR18] into a vocabulary of size 3052230522 with 128 tokens per sequence. Each sequence is mapped to a random binary sequence of the same length.

Learning to sort. Given an input sequence of letters, this task learns to sort the letters in alphabetical ordering (similar task have been studied before [FOŠ19]). Specifically, the model’s output for each input letter is used to determine the position of that letter in the predicted ordering. Each input sequence, of length 88, is created by sampling uniformly randomly, with replacement, from an alphabet of size 1010. The training and test sets consist of 1000 and 200 sequences, respectively. To ensure robustness with respect to hyperparameters, we experimented with a variety of settings (adjusting the model depth, number of heads, and the difficulty of the task by changing the alphabet size and sequence length) and observed consistent behavior.

Convex hull prediction. This task was inspired by the work of [VFJ15]. Given a sequence of NN points uniformly distributed in ×\times and shifted by a random bivariate standard normal, this task predicts the convex hull of these points. Specifically, for each point in the set, the model predicts whether it’s part of the convex hull. The training set consists of 10,00010,000 sequences of points in ×\times, each of length 1010.

In all three tasks, we report the test-set per-token label prediction accuracy as the evaluation metric.

We measure the effectiveness of individual paths by a ‘path disentanglement’ procedure that we apply at inference time: the procedure isolates the weights involved and the output of an individual path (PhLL⋯Ph11) X (Wh11⋯WhLL)(\bm{P}_{h_{L}}^{L}\cdots\bm{P}_{h_{1}}^{1})\,\bm{X}\,(\bm{W}_{h_{1}}^{1}\cdots\bm{W}_{h_{L}}^{L}) for any given sequence of heads h1,⋯ ,hL∈[H∪0]Lh_{1},\cdots,h_{L}\in[H\cup 0]^{L}. After the transformer has been successfully trained to solve each task (without modifications), we use this procedure to determine the output of a randomly sampled set of paths of a given length. We then evaluate the task performance based solely on the normalized sum of this subset of paths (rather than from all paths). Note that the training remains unaltered and uses all heads simultaneously, therefore ensuring that each path learns to its full effectiveness.

Figure 4 illustrates the resulting performance across all three tasks. We test different subset sizes and report the average and standard deviation of five repetitions. For reference, we also plot the accuracy of a naive classifier as well as of the entire trained model (i.e., before the path decomposition). As observed, short paths carry predictive power, with length-1 paths attaining accuracy above 0.8,0.6, and, 0.65 in the memorization, sorting, and convex hull tasks, respectively. On the other hand, the output of longer paths is not much better than a random guess (red horizontal lines). We note that, since there is a class imbalance in the convex hull task, we use a majority class predictor to obtain a random baseline. Though the difference in accuracy between short and long paths is less pronounced for the convex hull task, we observe that the variance of the long paths is significantly larger, rendering them not much better than a random guess. Length zero paths attain very small variance, but contain no useful information about the task (likely because they do not exploit global information).

The depths (LL), number of heads (HH), and hidden dimensions (dd) for the three models are: LL:6, HH:2, d:d:250 for memorization, LL:6, HH:2, dd:48 for sorting, and LL:6, HH:3, dd:84 for convex hull. It’s important to note that for all three tasks, while higher peak accuracies are attainable with increased model capacity and training time, our focus is to study the effects of path length on performance. Indeed, the trend for degenerating performance as path length increases stayed consistent across model sizes in all experiments.

The rapidly diminishing effectiveness of paths with respect to length indicates that the transformer relies almost exclusively on short paths. In other words, the transformer behaves like an ensemble of shallow networks. Furthermore, the results indicate that there is underutilized capacity in long paths, and suggest that one way to make them, and hence the transformer, more effective, is to prevent the long paths from losing rank.

Related works

Skip connections were first introduced in ResNets [HZRS16a], ever since, it has been used to facilitate optimization in deep networks [HZRS16b, VWB16b, BFL+18]. In particular, skip connections tackle the vanishing gradient problem, by allowing the gradient to flow bypass the skipped layers during backpropagation. The original motivation of using skip connections in transformers follow the same reasoning on facilitating optimization [VSP+17]. With the paths decomposition for transformers, we discover an additional surprising importance of skip connections: they prevent the transformer output from degenerating to rank one exponentially quickly with respect to network depth.

Veit et al. ([VWB16b]) introduced an analogous interpretation for residual networks as a collection of paths of varying lengths, and found that the length of the effective paths in deep residual networks are much shorter than the total network depth, due to the gradients used for parameter updates coming overwhelmingly from these short paths. Our finding suggests that SANs rely on short paths to avoid rank collapse. On the other hand, Daneshmand et al. [DKB+20] studied rank collapse in randomly initialized linear and ReLU networks and showed that batch normalization is an effective mitigation strategy.

Some recent works have approximated the attention matrix with low-rank factorizations [WLK+20, TBM+20] or kernel methods [KVPF20, CLD+20], to reduce the quadratic self-attention complexity. Our work is orthogonal to these works, by studying the rank of the network’s output (rather than of the attention matrix).

There have been other recent advances in understanding the theory behind transformers: [PMB19, DGV+19] proved Turing universality, [CLJ20b] provided necessary and sufficient conditions for attention to simulate convolution. A linearized form of self-attention was also found to exhibit a depth phase transition [LWS+20]; and the Lipschitz constant of self-attention was analyzed by [KPM20].

Perhaps the convergence to rank one of a path should come as no surprise: each path component contains row-stochastic matrices as a result of the softmax attention, and [AT77] showed the exponential convergence of products of stochastic matrices to rank one. While the intuition behind stochastic matrices driving convergence still applies, in deep attention networks these matrices interact in more complex ways than what classical analyses consider. As we show, because of these interactions the rank collapses much faster than what would be expected based on classical analyses (cubic vs linear rate).

Conclusion

This work exposes competing forces over rank collapse in self-attention networks, namely self-attention vs skip connections and MLPs. In the process, we develop a path decomposition for SANs, which modularizes the study of self-attention and is of independent interest to additional applications. These results open the door for many exciting future directions. For instance, how can one leverage the token-uniformity inductive bias revealed to design more effective networks, perhaps better at utilizing long paths? What are some practical implications for width-depth trade-off? How do we prove meaningful lower bounds of residue convergence for transformers? Answering these questions has broad implications in advancing the state of the art in deep learning.

Acknowledgements. Andreas Loukas would like to thank the Swiss National Science Foundation for supporting him in the context of the project “Deep Learning for Graph-Structured Data” (grant number PZ00P2 179981). Jean-Baptiste Cordonnier is supported by the Swiss Data Science Center (SDSC).

References

Appendix A Deferred Proofs

We build our argument step by step, by first considering a single-head self-attention layer in §A.1 and then moving to deeper networks with single and multiple heads in §A.3 and §A.4. The results are extended to take into account skip connections and MLPs in §A.5 and §A.6

We consider a single-head self-attention layer:

We focus in particular on how the residual changes. As discussed previously, the value bias can be safely ignored since it does not contribute to the residual.

with γ\gamma selected such that max⁡i,j,j′∣Aij−Aij′∣ ∑imax⁡j,j′∣Aij−Aij′∣≤γmax⁡j,j′∑i∣Aij−Aij′∣\sqrt{\max_{i,j,j^{\prime}}|A_{ij}-A_{ij^{\prime}}|\,\sum_{i}\max_{j,j^{\prime}}|A_{ij}-A_{ij^{\prime}}|}\leq\gamma\max_{j,j^{\prime}}\sum_{i}|A_{ij}-A_{ij^{\prime}}| and ∣Eij−Eij′∣≤1.256|E_{ij}-E_{ij^{\prime}}|\leq 1.256 with E=res(X)WQKdqkres(X)⊤\bm{E}=\text{res}(\bm{X})\frac{\bm{W}_{QK}}{\sqrt{d_{\textit{qk}}}}\text{res}(\bm{X})^{\top}.

The unscaled attention scores are computed as follows,

and following [CLJ20a], we can use the softmax shift invariance property to prune the terms constant over the columns and obtain,

with WQK=WQWK⊤\bm{W}_{QK}=\bm{W}_{Q}\bm{W}_{K}^{\top} and bQK=WKbQ\bm{b}_{QK}=\bm{W}_{K}\bm{b}_{Q}.

We use the shorthand notation R:=res(X)\bm{R}:=\text{res}(\bm{X}) and R′:=res(X′)\bm{R}^{\prime}:=\text{res}(\bm{X}^{\prime}).

Using the shift-invariance property of the softmax operator, the first term above can be safely ignored since it is constant across columns. We therefore have that

where we have set r:=RWQK⊤dqkx+RbQKdqk\bm{r}:=\bm{R}\frac{\bm{W}_{QK}^{\top}}{\sqrt{d_{\textit{qk}}}}\bm{x}+\bm{R}\frac{\bm{b}_{QK}}{\sqrt{d_{\textit{qk}}}}.

where the inequality above is entry-wise and follows from Lemma A.3 whenever ∣Eij−Eij′∣≤1.256|E_{ij}-E_{ij^{\prime}}|\leq 1.256. Similarly PX≥1(x⊤+softmax(r)⊤R)−2D 1 softmax(r)⊤R\bm{P}\bm{X}\geq\text{{1}}(\bm{x}^{\top}+\text{softmax}(\bm{r})^{\top}\bm{R})-2\bm{D}\,\text{{1}}\,\text{softmax}(\bm{r})^{\top}\bm{R}, where we again invoke Lemma A.3.

Therefore, the (entry-wise) distance of the output of the self-attention layer SA(X)=PXWVSA(\bm{X})=\bm{P}\bm{X}\bm{W}_{V} from being constant across tokens is at most:

where r′=(x+R⊤softmax(r))WV\bm{r}^{\prime}=(\bm{x}+\bm{R}^{\top}\text{softmax}(\bm{r}))\bm{W}_{V}.

where the last step is due to ∥softmax(r)∥1=1\|\text{softmax}(\bm{r})\|_{1}=1 and ∥AB∥1≤∥A∥1∥B∥1\|\bm{A}\bm{B}\|_{1}\leq\|\bm{A}\|_{1}\|\bm{B}\|_{1}, implying ∥SA(X)−1(r′)⊤∥1≤2∥D1∥1 ∥R∥1∥WV∥1\|SA(\bm{X})-\text{{1}}(\bm{r}^{\prime})^{\top}\|_{1}\leq 2\|\bm{D}\text{{1}}\|_{1}\,\|\bm{R}\|_{1}\|\bm{W}_{V}\|_{1}.

Moreover, by the definition of D\bm{D} as in Lemma A.3 and under the current Lemma’s definition, we have that

A.2 Multiple-heads and single-layer

In the setting of Lemma A.1, the residual of the output of a HH-heads attention layer abides to:

where ∥WQK,h∥1∥Wh∥1,∞≤β\|\bm{W}_{QK,h}\|_{1}\|\bm{W}_{h}\|_{1,\infty}\leq\beta for all heads h∈[H]h\in[H].

The output of a multi-head attention layer is

where Wh:=WV,hWO,h\bm{W}_{h}:=\bm{W}_{V,h}\bm{W}_{O,h} as in the main text and Ph\bm{P}_{h} is computed using the heads parameters WQK,h\bm{W}_{QK,h} and bQK,h\bm{b}_{QK,h}. The proof proceeds similarly to Section A.2 until eq. 11,

where r′′=∑h(x+R⊤softmax(rh))Wh\bm{r}^{\prime\prime}=\sum_{h}(\bm{x}+\bm{R}^{\top}\text{softmax}(\bm{r}_{h}))\bm{W}_{h}.

A.3 Single-head and multiple-layers

We next consider how the residual changes after LL layers of the form: Xl=SA1l(Xl−1).\bm{X}^{l}=\text{SA}^{l}_{1}(\bm{X}^{l-1}).

In the setting of Lemma A.1, for any single-head SAN consisting of LL layers with ∥WQK,1l∥1≤β\|\bm{W}_{QK,1}^{l}\|_{1}\leq\beta for every l∈[L]l\in[L], the residual is bounded by

which amounts to a doubly exponential convergence to a rank-1 matrix.

Unfolding the recursion backwards from the last layer to the first and applying Lemma A.1 we obtain:

A.4 Multiple-head and multiple-layers

In the setting of Lemma A.1, consider a depth-LL SAN with HH heads per layer. Fix ∥WQK,hl∥1∥Whl∥1,∞≤β\|\bm{W}_{QK,h}^{l}\|_{1}\|\bm{W}_{h}^{l}\|_{1,\infty}\leq\beta for all h∈[H]h\in[H] and l∈[L]l\in[L]. The output residual is bounded by

which indicates that the output convergences to a rank-1 matrix doubly exponentialy.

The proof procceeds recursively as for Theorem 2.2 in the single head case but using the bound on single-layer multi-heads residuals from Lemma A.2. ∎

A.5 SAN with skip connections

As noted in the main text, a lower bound on the residual better aligns with practice, where SANs with skip connections do not suffer rank collapse. For consistency with the other analyses and as one way to illustrate residual growth, we provide a (vacuously large) upper bound on the residual for SANs with skip connections.

In the setting of Lemma A.1, consider a depth-LL SAN with HH heads per layer and skip connections. Fix ∥WQK,hl∥1∥Whl∥1,∞≤β\|\bm{W}_{QK,h}^{l}\|_{1}\|\bm{W}_{h}^{l}\|_{1,\infty}\leq\beta for all heads h∈[H]h\in[H] and layers l∈[L]l\in[L]. The output residual is bounded by

For a SAN with skip connections, the residual bound for a single-head single-layer SAN from lemma A.1 now becomes:

To obtain a multi-layer bound, we unfold the recursion backwards.

Let us consider a single head model first and fix ∥WQK,hl∥1∥Whl∥1,∞≤β\|\bm{W}_{QK,h}^{l}\|_{1}\|\bm{W}_{h}^{l}\|_{1,\infty}\leq\beta for all l∈[L]l\in[L]. We have that:

Now we unroll this bound across layers to write it in terms of res(X)\text{res}(\bm{X}). At the kthk^{th} step of unrolling, the max is one of the two terms in Eq 24: either 4 γ βdqk ∥res(XL−k)∥1,∞3\frac{4\,\gamma\,\beta}{\sqrt{d_{\textit{qk}}}}\,\|\text{res}(\bm{X}^{L-k})\|_{1,\infty}^{3} or ∥res(XL−k)∥1,∞\|\text{res}(\bm{X}^{L-k})\|_{1,\infty}, i.e. we make a binary choice. Thus unrolling through all LL layers corresponds to a path from the root to the maximum leaf in a depth-LL complete binary tree. Each leaf has the form (8 γ βdqk)3l−12 23l(L−l)∥res(X)∥1,∞3l\left(\frac{8\,\gamma\,\beta}{\sqrt{d_{\textit{qk}}}}\right)^{\frac{3^{l}-1}{2}}\,2^{3^{l}(L-l)}\|\text{res}(\bm{X})\|_{1,\infty}^{3^{l}}, where ll indicates the number of times the term 4 γ βdqk ∥res(XL−k)∥1,∞3\frac{4\,\gamma\,\beta}{\sqrt{d_{\textit{qk}}}}\,\|\text{res}(\bm{X}^{L-k})\|_{1,\infty}^{3} is chosen as the max. Note the ordering of these choices does not matter, only the number of times a term is chosen. Consequently, the residual bound is the maximum amongst such leaf terms:

We now apply this bound to HH heads, we use Lemma A.2, which for a single layer gives:

Therefore, accounting for the factor of HH in above, we obtain a residual bound for a depth-LL width-HH SAN with skip connections:

A.6 SAN with MLP

We now study how using an MLP affects the residual. Recall we focus on SANs with layers written as

Note that, to keep the notation compact, we use flf_{l} to encompass both the MLP as well as the output bias.

The proof proceeds the same way as in §A.1. For clarity, we point out the differences with proof in §A.1 without repeating details that remain the same.

In the setting of Lemma A.1, consider a depth-LL and width-HH SAN with MLP. Moreover, let ∥WQK,hl∥1∥Whl∥1,∞≤β\|\bm{W}_{QK,h}^{l}\|_{1}\|\bm{W}_{h}^{l}\|_{1,\infty}\leq\beta for all h∈[H]h\in[H] and l∈[L]l\in[L] and fix λl,1,∞≤λ\lambda_{l,1,\infty}\leq\lambda. We then have that

With an MLP as formulated in Eq 25, we have Wh:=WVWO\bm{W}_{h}:=\bm{W}_{V}\bm{W}_{O} in place of just the value weight WV\bm{W}_{V}, as defined in the main text. As before, let R\bm{R} denote res(X)\text{res}(\bm{X}).

The proof proceeds the same way as in Lemma A.1, until Eq 11, where we handle the multi-head case the same way as in Eq A.2 to obtain the entrywise inequality:

We now use the fact that f(1r′⊤)f(\text{{1}}r^{\prime\top}) also takes the form 1r′′⊤\text{{1}}r^{\prime\prime\top} for some vector r′′r^{\prime\prime}. Indeed, ff encompasses weight matrix multiplications, bias addition, and entrywise nonlinearities, all of which preserve the fact that f(1r′⊤)f(\text{{1}}r^{\prime\top}) is constant across rows. Therefore,

Subsequently, just like for the single-head single-layer proof, we bound ∥Dh 1 softmax(rh)⊤RWh∥p\|\bm{D}_{h}\,\text{{1}}\,\text{softmax}(\bm{r}_{h})^{\top}\bm{R}\bm{W}_{h}\|_{p} in the above by

As we have shown before, ∥Dh1∥1,∞\|\bm{D}_{h}\text{{1}}\|_{1,\infty} can be bounded above by 2γ dqk ∥R∥1∥WQK,h∥1∥R∥∞\frac{2\gamma\,}{\sqrt{d_{\textit{qk}}}}\,\|\bm{R}\|_{1}\|\bm{W}_{QK,h}\|_{1}\|\bm{R}\|_{\infty}. Applying this to both Eq 28 and Eq 29, and combining the two as in Lemma A.1, yields the bound:

Finally, we recursively unroll the bound across layers to obtain a residual bound in terms of res(X)\text{res}(\bm{X}):

A.7 A technical lemma

with the diagonal matrix D\bm{D} having Dii=max⁡j,j′∣δi⊤E(δj−δj′)∣D_{ii}=\max_{j,j^{\prime}}|\bm{\delta}_{i}^{\top}\bm{E}(\bm{\delta}_{j}-\bm{\delta}_{j^{\prime}})| and the inequality taken entry-wise.

Let us start by the definition of the row-stochastic matrix:

The above, implies that for every i,ji,j we have:

which holds for ∣Eij−Eij′∣≤1.256|E_{ij}-E_{ij^{\prime}}|\leq 1.256. Notice also that

both of which are at most max⁡j′∣δi⊤E(δj−δj′)∣\max_{j^{\prime}}|\bm{\delta}_{i}^{\top}\bm{E}(\bm{\delta}_{j}-\bm{\delta}_{j^{\prime}})|, from which the claim follows. ∎

Appendix B Additional results

As we saw in §2.1, transformers can be viewed as an interdependent ensemble of simpler networks (or paths) each of different depth (or length). Aiming to gain more insight about the ensemble structure in practice, Fig 5 visualizes the path length distribution in various commonly-used architectures.

Based on the exponential decay of path effectiveness result, we hypothesize that models that focus overwhelmingly on long paths are less efficient than models with a more diverse path distribution. The long-paths models are furthermore likely to be less robust, as they require larger MLP Lipschitz constants to counteract the token-uniformity inductive bias caused by self-attention, as described in §3. It is perhaps no coincidence that the intentionally more efficient models, such as DistilBert or MobileBert, have some of the most diverse path distributions; and that for the most extreme long-paths-focused model, GPT3, studies found that its model size can be reduced by several orders of magnitude and achieve similar performance [SS20]. We leave these exciting directions for future work.