Signal Propagation in Transformers: Theoretical Perspectives and the Role of Rank Collapse
Lorenzo Noci, Sotiris Anagnostidis, Luca Biggio, Antonio Orvieto, Sidak Pal Singh, Aurelien Lucchi
Introduction
Since its first appearance in Vaswani et al. (2017), the Transformer architecture has revolutionized the field of Natural Language Processing (NLP), achieving remarkable success in tasks such as text classification (Yang et al., 2019), machine translation (Conneau and Lample, 2019), reading comprehension (Brown et al., 2020) and question answering (Raffel et al., 2019) among others. Recent efforts have effectively extended its applicability to computer vision (Dosovitskiy et al., 2020) and other domains (Baevski et al., 2020; Huang et al., 2018; Biggio et al., 2021; Polu et al., 2022), further popularizing it outside NLP.
The Transformer operates on inputs comprising a sequence of tokens. At its core, it relies on stacked attention layers, which compute a measure of relevance for the whole sequence by assigning token-wise importance weights — obtained by matrix multiplication of the queries and keys, and finally normalized with the softmax function. The output of an attention layer is then a linear combination of the importance weights and the so-called values. Then, the architecture includes fully-connected sub-layers, residual connections (He et al., 2016), and layer normalization (LN), as illustrated in Fig. 2.
In the absence of residual connections, Dong et al. (2021) proved that at initialization the rank of the sequence representation collapses doubly exponentially with depth, and both layer normalization and fully connected layers can only partially alleviate the speed of degeneracy. Under rank collapse, the model does not distinguish between representations of different tokens, which are perfectly aligned in feature space at initialization. Crucially, the precise implications of rank collapse in Transformers are not fully understood.
In this paper, we show that a high alignment of the tokens’ representations at initialization — corresponding to rank collapse in the extreme case of perfect alignment — affects training by causing vanishingly small gradients of the queries and keys’ parameter matrices. This problem severely diminishes the capabilities of the model to learn meaningful attention weights and is further exacerbated in very deep networks, where the rank deficiency — and hence the vanishing gradient problem of the queries and keys — affects several layers (see Fig. 1). In order to shed light on this problem, we take inspiration from the flourishing literature on signal propagation in random networks and start our analysis by computing the expected gradients of an attention layer with respect to the queries, keys, and values, which leads to Theorem 3.1 on the vanishing gradients for the queries and keys. From here, we pursue two different directions.
Firstly, we investigate under which conditions rank collapse can be avoided by studying the evolution of the input sequence in a Transformer at initialization. Our theory reveals that a depth-dependent scaling of the residual branches, beyond stabilizing the norm of the activations at initialization, also approximately preserves the cosine of the angle between tokens, and hence also stabilizes the rank of the propagating sequence. We show that this holds even in the infinite-depth limit.
Secondly, we illustrate that there are factors, other than the average tokens’ correlation, that affect differently the gradient norm of the queries and keys compared to the values. In particular, the propagating sequence’s squared norm has a linear dependence in the values, while a cubic one in the queries and keys, justifying the use of layer normalization. We also highlight a different dependence on the embedding dimension and the length of the input sequence, implying that the gradient norm of a subset of parameters can potentially be of different orders of magnitude, as empirically hinted by previous works (Liu et al., 2020). Our analysis brings to light fundamental issues in the signal propagation in Transformers, opening the way for new, well-founded and motivated approaches to improve optimization in these models.
Background
A Transformer architecture consists of stacked attention blocks, as show in Fig. 2. Layer normalization is usually applied token-wise either after the residual connections or to the inputs of the self-attention and position-wise feed-forward sub-layers, leading to the POST-LN (Vaswani et al., 2017) and PRE-LN (Wang et al., 2019; Xiong et al., 2020) variants respectively.
Theoretical Results
2 Forward Signal Propagation and the Importance of Scaling the Residual Branches
We now turn our attention to the study of the influence of skip connections in transformers. Dong et al. (2021) showed that simply adding skip connections prevents rank collapse. Somewhat surprisingly, we show that while the claim holds for any finite depth, the average angle between different tokens quickly increases with just a few layers, and as a Transformer can still lose rank unless the residual branches are adequately initialized. As Dong et al. (2021) showed that layer normalization does not avoid rank collapse, we omit it in our analysis. Firstly, we introduce two lemmas on the propagation of inner products (Lemma 3.2) and the norm (Lemma 3.3) of the tokens’ representations.
The previous Lemma provides theoretical justification that scaling the residual branches by setting the alpha parameters to be allows both the norm of the propagating input and the inner products between different tokens to be approximately preserved. Hence, the information contained in the input is not lost, even in the infinite depth limit.
We are now ready to formalize the influence of the -scaling on the correlation between tokens’ representations by stating Theorem 3.2.
3 Dependence on the Angle between Tokens and the Input Norm
By making the additional assumption that the norm and the correlation propagate independently, the respective norm for the queries — and symmetrically the keys — (Eq. (7)) reduces to:
In Appendix A.2 we provide a rigorous proof, that relies on Isserlis theorem (Isserlis, 1918) to compute higher-order moments. The above expressions reveal the different dependencies on four main actors, that we inspect separately here. The gradients of the queries depend via a cubic function on the variance of the input, , compared to a linear for the values. This provides an additional interpretation of the successful use of layer normalization, as in Xiong et al. (2020), either in the POST-LN or PRE-LN format, that standardizes the input variance to the value .
Next, we emphasize the dependence on the correlation between the tokens, also illustrated in Fig. 3. Importantly, note how the queries/keys have opposite monotonic functional dependence with respect to compared to the values. As revealed by Theorem 3.2 and Fig. 3 (center), inappropriate scaling of the residual branches can already lead to this phenomenon even in a relatively shallow network.
Finally, Eq. (17) and (18) reveal a different scaling in terms of the embedding size and the sequence length due to the self-attention operation itself. We hope that the identification of the different dependencies in the gradients of the parameters will inspire a new line of works aimed at solving some of the difficulties in training Transformers.
4 Are Adaptive Methods really needed for training Transformers?
A direct consequence of our analysis is that allows controlling the magnitude of the gradients for the queries and keys’ parameters.
We evaluate our proposal, consisting of residual scaling and the aforementioned inverse temperature parameters, on the widely used IWSLT14 German-to-English (De-En) benchmark translation task. All details regarding the experimental setup and the choice of inverse temperature used are provided in Appendix C. We train a Transformer encoder-decoder of varying depth with SGD, after removing all normalization layers and adequately initializing the residual connections. For our training with SGD, we avoid using any learning rate warm-up, as commonly done for Adam, and instead use a step-scheduler to decrease the learning rate at 40% and 80% of training. We compare against the following methods that make use of Adam; POST-LN and PRE-LN refer to the aforementioned alternatives to apply layer normalization. We also compare against other successful techniques that rely on specific initializations to avoid layer normalization, such as ReZero (Bachlechner et al., 2021) and T-Fixup (Zhang et al., 2019). We report the average BLEU score (Papineni et al., 2002) across 5 runs in Fig. 7 and Table 7.
Our proposed method considerably improves training with SGD, keeping up and in some cases surpassing any results achieved by the Adam optimizer. We are also able to train deeper networks without the use of layer normalization. We leave for future work to further investigate modifications or alternatives to the self-attention operation.
Related Work
Our work builds upon the rich literature on forward and backward signal propagation in random neural networks (Poole et al., 2016; Schoenholz et al., 2017; Xiao et al., 2018; Pennington et al., 2017; Orvieto et al., 2021; Noci et al., 2021; Zavatone-Veth and Pehlevan, 2021). The scaling scheme has been investigated in the literature for the stabilization of residual networks (Hanin and Rolnick, 2018; Arpit et al., 2019; Allen-Zhu et al., 2019; Hayou et al., 2021).
Our work draws inspiration from a series of recent works studying the rank of the representations of random feed-forward neural networks at initialization (Daneshmand et al., 2020b, a). In the context of Transformers, Dong et al. (2021) has recently identified the rank collapse issue object of study of the present work. Thanks to our analysis of the backward pass, we are able to demonstrate that rank collapse in Transformer architectures leads to vanishingly small gradients of queries and keys, thereby preventing effective training and allowing us to complete the analysis of (Dong et al., 2021).
Among the architectural components in Transformers, layer normalization is, arguably, one of the most important – and debated – ones (Chen et al., 2018; Wang et al., 2019; Nguyen et al., 2010; Xiong et al., 2020). In the original architecture (Vaswani et al., 2017), layer normalization is used to stabilize the forward pass by reducing the variance of the inputs to the following sublayer. Our analysis of the forward pass shows that its inclusion is not strictly necessary for the purpose of controlling the norm of the representations. For a theoretical analysis of signal propagation in the presence of layer norm, we refer the reader to Xiong et al. (2020).
Additionally, our theoretical study of the backward pass provides a rigorous explanation of the empirically observed discrepancy between the magnitude of the gradients of the queries and the values, which Liu et al. (2020) hypothesize to be one of the causes of the success of adaptive methods in training Transformers (Liu et al., 2019; Zhang et al., 2020; Huang et al., 2020).
Finally, properly rescaled residual connections have been found to be beneficial for training Transformers by a number of recent research works (Zhang et al., 2019; Bachlechner et al., 2021; Wang et al., 2022). However, none of these studies characterize the impact of skip connections on rank propagation, while our analysis suggests a theoretically-grounded way to stabilize it.
Conclusions and Future Work
In this paper, we showed how, at initialization, rank collapse and more generally high correlation in the tokens, causes vanishing gradients of the queries and keys of a Transformer architecture. While residual connections help mitigate rank collapse at finite depth, we showed that they alone cannot prevent high alignments of the tokens’ representations — unless properly scaled by a -factor. Finally, we have also discovered counter-intuitive dependencies on the variance of the input, embedding size, and sequence length, potentially causing large differences between the gradients of queries/keys compared to the values’ parameters. Hence, we conclude that one of the strengths of Transformers lies in their carefully designed architecture together with an adequate initialization. Finally, we gave preliminary evidence that one of the factors contributing to the higher efficacy of Adam compared to SGD in training Transformers arises from the disproportionate magnitude of gradients as postulated by our theory. Nonetheless, other factors might further accentuate the difference between these two algorithms during training, leaving the door open for further research regarding the benefits of adaptive optimization methods with Transformers.
Acknowledgements
We thank our colleague Jonas Kohler who provided insights and expertise that greatly assisted the research.
References
Appendix
Recall the defining equations of a Transformer:
where the self attention layers are defined as follows:
Initialization: Recall that we initialize our weights with the so called “Xavier” [Glorot and Bengio, 2010b] or “He” [He et al., 2015] initialization: each weight is sampled independently from a distribution with zero-mean and variance for the values and feedforward weights and for the queries and keys.
Kronecker Delta: we introduce the Kronecker Delta notation
and, similarly: .
Notation: in this section, we also adopt the following shorthand notation for the argument of the softmax . We will first compute the gradients with respect to values, queries and input (note that the gradients of the keys have the form as the queries, hence we omit the derivation). Recall that for a matrix , we use to indicate its -th row. Finally, we indicate with the matrix with all ones, with the columns vector with all ones, and with the -dimensional identity matrix.
In this section, we now look at the proofs of Lemma 3.1 and Theorem 3.1. We will first introduce our notation for the gradients, as well as some useful properties of the Kronecker product.
For the gradients, we avoid directly working with tensors by vectorizing the matrices in a row-wise fashion () and arranging the Jacobian in the numerator layout. More formally,
Alongside this, we use the following rule ( is the Kronecker product):
For the proof of this rule, we refer to Singh et al. , and to Magnus and Neudecker for a complete introduction to matrix calculus. We will also use the following well-known properties of the Kronecker product.
In Lemma A.2 and Lemma A.3 we compute the gradients with respect to the queries, values and , respectively. Then we use these results to prove Lemma 3.1 by computing the expectation of the Frobenius norms.
Lemma A.2 (Gradients of Self Attention for parameter matrices). The gradients of the self attention layer defined in Eq. (1) have the following form: where the gradients of the softmax with respect to its inputs are as follows: \frac{\partial{\bf A}}{\partial{\bf M}}=\operatorname{blockdiag}\Bigg{(}\dfrac{\partial{\bf A}_{i}}{\partial{\bf M}_{i}^{\top}}\Bigg{)} (23) and where with being the -th row of in column vector format. Finally, note that under the uniform-attention assumption, Eq. (23) simplifies to: (24) Proof. Let’s start with the simple case of the values’ weights . Using the rule in Eq. (20), it is immediate that:
For the queries, a simple application of the chain rule and then again Eq. (20) gives:
which is the desired results. Finally, for the gradients of the softmax note that:
By writing the above expression in the matrix notation described above, we obtain the desired result. More specifically, the block diagonal structure is given from the term which stems from the fact that the softmax is applied row-wise. ∎
Lemma A.3 (Gradients of Self Attention with respect to the Embedding matrix). The gradients of the self attention layer with respect to the embedding matrix defined in Eq. (1) have the following form (25) where the gradients of the softmax with respect to its inputs are denoted by as before. Proof. Remember that we defined . Alongside with our previous shorthands , , let us define the remaining as a matrix , so that . Both and are functions of . So the matrix differential can be written as:
Next, we use the matrix differential and then the identification theorem of matrix derivatives to compute the matrix gradient
and . Therefore, for rowwise vectorization, we have a similar result:
where in the last line we used the fact the commutation is a permutation matrix, so . Thus, we get the required matrix derivative as follows:
Next, we will use a property of commutation matrix to make things simpler (Theorem 7.9, Magnus and Neudecker ):
Plugging this into the above Eq. (29), we get:
Gradient with respect to the values matrix.
Gradients with respect to the queries/keys matrix.
First, recall the expression for the gradient of the softmax under the uniform-attention assumption (Eq. (24)):
Hence, we can rewrite the expression of Lemma A.2 for the gradients of the queries as:
where in the last step we have used twice the property of the Kronecker product in Eq. (22) of Lemma A.1.
where we have used the property on the trace of the Kronecker product (Lemma A.1, Eq. (21)). Note that if we are conditioning on , then we only have to take the expectation of the last term with respect to the weights and . Let us call for notation simplicity.
Plugging in the values of and under the uniform-attention assumption into Eq. (25) gives rise to the following:
Let’s refer to the matrices on the right-hand side as respectively. We compute the expected squared Frobenius norm of these as follows:
where in the second line we have taken the expectation inside and used the fact that , being a commutation matrix, is orthogonal. Then, by simple properties of Kronecker product and cyclic property of trace, we have the result, which is the same as that for .
See 3.1 Before starting the proof, it is interesting to note that, even though the gradients of queries and keys vanish in the rank collapse regime (i.e. ), the gradient with respect to the values and the input does not (see Theorem 3.1). From this simple remark, we can conclude that, even in the rank collapse regime, information still propagates in the backward pass. In Section 3.4 (main paper), we show that even if gradients effectively propagate, the phenomenon studied in this theorem still greatly affects training.
By using the chain rule and the fact that for two matrixes we have that , we can upper bound the gradient as:
A.2 Gradient Analysis of Section 3.3
Throughout this section we assume that between every pair of tokens, the same dimension is a zero-mean Gaussian random variable with the same correlation, meaning that
As we will deal with the computation of 4-th order moments of correlated Gaussian random variables, we will make use of Isserlis theorem [Isserlis, 1918]:
Now we can prove Eq. (17) , which we re-state here:
Now, we have 2 cases: if , which gives equal terms, we need to compute
where is any tuple with and we used uncorrelation of different dimensions. Otherwise, we get each equal to
where we leveraged the following direct calculations:
Finally, similar computations also lead to the last term. To follow the calculations, we invite the reader to draw the matrix and to hide the columns over which summations are not performed:
We make use of Isserlis theorem, stating that:
By using our independence assumptions, we get:
Let’s study it term by term. We will also use , and so .
First term: we have that which is equal to (omitting the constant ):
Second term: we have that which is equal to (omitting the constant ):
Third term: we have that which is equal to (omitting the constant ):
Plugging in the values of A and B we get:
and finally assuming Xavier initialization
A.3 Forward Pass: Proofs of Lemma 3.2 and 3.3
First, we characterize the evolution of the correlations between tokens with depth, under the assumptions of Theorem 3.3, namely uniform-attention assumption, and the adoption of a linear activation.
Note that under the uniform-attention assumption:
Hence, using the fact that the weights are i.i.d with variance :
where in the last step we have unrolled the recursion until the input layer.
For the limit as , simply note that:
Now we are ready to re-state and prove Lemma 3.3.
See 3.3 Proof. The proof is in the same spirit as Lemma 3.2 but slightly more involved. Again, using Lemma A.5 in both the skip connections of the Transfomer architecture. Therefore, using Lemma A.5 (skip), Lemma A.4 (linear) and Lemma A.6 (attention):
Let . For the latter term we have that:
The final results as stated in the theorem hold because of the following:
Remark: note that . ∎
A.4 Proof of Theorem 3.2: Correlations are Preserved under Residual Scaling
A.5 Motivation for Assumption 3.1
We motivate here the following assumption, stated in the main paper. This assumption is crucial to compute expectations involving the softmax function.
We first show that this assumption holds when taking to infinity, keeping fixed.
The following classical result implies almost sure convergence of the softmax matrix as .
Let be a sequence of random variables. If for any
then converges to almost surelyThat is, for almost every (i.e. with probability one)..
Borel Cantelli then directly yields almost sure convergence of to as . Next, note that both and are continuous functions of , hence we can apply standard continuity event-per-event. For almost every ,
Hence almost surely. This can also be seen as a simple application of the continuous mapping theorem. The same reasoning yields almost sure convergence of
to the corresponding limiting quantity. ∎
Appendix B Additional Results
We present some additional results on the propagation of the norm and the correlations in Figure 9. In particular, we empirically show that, with an adequate depth-dependent residual scaling, the norm and the correlation are stabilized, even for very deep networks. Furthermore, we demonstrate the propagation of the correlation and the gradient norms for the PRE-LN configuration in Figure 10. As also hinted in the main text, in Figure 4, the increase in correlation with depth for PRE-LN is much less wild. This also results in better stabilized gradients for the queries and keys’ parameters. We also observe the opposite trend for the gradients of the values, in relation to the POST-LN case in Figure 3. We speculate that this different dependence, along with the better preserved correlation, is the main reason PRE-LN configured Transformers have been shown to scale better with depth. We plan to investigate this dependence more in future work.
B.2 Further Empirical Assessment of Assumption 3.1
Here, we empirically test the accuracy and limitations of the uniform-attention assumption.
For the empirical verification of Assumption 3.1 in the forward pass analysis, we plot the density of the norm of the representations for only-encoder Transformers of increasing depth. The results are shown in Fig 11. Note that when the standard deviation of the input is set to , then the uniform-attention assumption provide an excellent approximation to the common Xavier-initialization. On the contrary, we observe a deviation when the standard deviation of the input is increased. Also, note how as the depth increases, the distribution becomes more heavy-tailed. This heavy-tailedness was recently formally shown for standard MLPs with and without ReLU activation [Noci et al., 2021, Zavatone-Veth and Pehlevan, 2021].
For the verification of the assumption in the backward pass, we additionally show in Fig. 12 how the norm of the gradients w.r.t queries and keys depends on the hidden dimension, the sequence length, the input correlation and the input variance. Ground-truth gradients are calculated with automatic differentiation, and they are compared with our theoretical results based on Assumption 3.1. As shown in Fig.12, our theoretical predictions show a very good agreement with the true gradients. Again, we notice that the smaller the values of the input standard deviation the tighter the agreement of the theory with the simulations. Intuitively, a higher input variance causes the argument of the softmax to have a large range of values. This in turn causes a deviation from the uniform distribution (i.e. maximum entropy), towards the distribution of minimum entropy (a Delta Dirac, corresponding to attending to only one token).
B.3 Empirical Verification of the Gradient Analysis of Section 3.3
Finally, in Figures 13 and 14 we show the dependence of the norm of the gradients for the keys and values based on the parameters of the architecture and the task-specific parameters. Figure 13 illustrates the true dependence and Figure 14 the one expected by the theory based on our assumptions. In short, the main takeaways are the following.
As the correlation between the tokens increases (-axis in the global plot), the norm of the gradients of the queries quickly diminishes compared to the one of the values.
The dependence on the variance of the input is different (-axis in the global plot), being linear for the values and cubic for the queries. This highlights the importance of a stabilized forward pass and provides another explanation regarding the successful use of layer norm in Transformers.
The dependence on (-axis in each subplot) and (-axis in each subplot) is more complicated, also being a function of the correlation (compare the first column where to the rest).
Appendix C Experimental Setup
Here we provide more details regarding the experimental setup.
In Figure 5, we focus on a toy example where the task is to reverse a sequence of tokens. More specifically, given a sequence of numbers in the range , we predict the same tokens in the inverted order. We use an embedding layer of size 16, initializes with variance 1, and sinusoidal positional encodings to initially embed the input. We use a 5-layer POST-LN Transformer encoder model, with a single head attention operation and a two-layer feed-forward layer with a ReLU nonlinearity. We use residual scaling in this case equal to . We train using Adam with betas parameters , learning rate and weight decay 0.
C.2 Translation Task
Introducing an inverse temperature scaling inside the softmax, modifies the attention operation to
Then the gradient of the queries and keys parameters are directly scaled by this temperature value, following the same proof as for Eq. (7). We choose a temperature value of to match the gradient norms of the values and queries as in Equations. (17) and (18). Doing so, we assume a constant small correlation between tokens (also empirically verified in Fig. 15) and set the sequence length to the average found in our training dataset. Due to instabilities in training, we use warm-up on this temperature value. In short:
with ‘’ and ‘step’ the current training step.
We base our implementation on fairseq [Ott et al., 2019]. For the hyperparameter configuration, we mostly rely on the extensive search already done in fairseq [Ott et al., 2019] and Liu et al. . The final used parameters are exhibited in Table 1. For the final evaluation, we use the best-performing model on the left-out validation set. We apply weight decay as in Loshchilov and Hutter for both SGD and Adam.
Finally, in Figure 15 we display the evolution of correlations, residual scaling, and norm of the activations, with depth, for our best trained model. The residual scaling are trainable parameters. This enables them to weight differently the residual branches if deemed necessary. Although these values increase during training, the correlation between the tokens does not significantly increase, which as implied by our main results, allows efficient propagation of the gradients. The norm of the propagated forward signal tends to slightly increase with depth.