Deep Transformers without Shortcuts: Modifying Self-attention for Faithful Signal Propagation
Bobby He, James Martens, Guodong Zhang, Aleksandar Botev, Andrew Brock, Samuel L Smith, Yee Whye Teh
Introduction
Despite numerous impressive successes, the practice of training deep neural networks (DNNs) has progressed to a large extent independently of theoretical justification. Most successful modern DNN architectures rely on particular arrangements of skip connections and normalisation layers, but a general principle for how to use these components in new architectures (assuming they are even applicable) remains unknown, and their roles in existing ones are still not completely understood.
The residual architecture, arguably the most popular and successful of these, was first developed in the context of convolutional networks (CNNs) (He et al., 2016), and later in self-attention networks yielding the ubiquitous transformer architecture (Vaswani et al., 2017). One proposed explanation for the success of residual architectures is that they have superior signal propagation compared to vanilla DNNs (e.g. Balduzzi et al., 2017; Xiao et al., 2018; Hayou et al., 2019; De & Smith, 2020; Martens et al., 2021), where signal propagation refers to the transmission of geometric information through the layers of a DNN, as represented by a kernel function (Daniely et al., 2016; Poole et al., 2016; Schoenholz et al., 2017).
Recently, using signal propagation principles to train DNNs at high depths, without the skip connections and/or normalisation layers found in residual architectures, has become an area of interest in the community. The reasons are two-fold. First, it would validate the signal propagation hypothesis for the effectiveness of residual architectures, thus clarifying our understanding of DNN trainability. And second, it could lead to general principles and techniques for achieving trainability in DNNs beyond the residual paradigm, with the potential for improved or more efficient architectures.
For CNNs, Xiao et al. (2018) showed that improved signal propagation from better initialisation enables very deep vanilla networks to be effectively trained, although at significantly reduced speeds compared to residual networks. Martens et al. (2021) later proposed Deep Kernel Shaping (DKS) which uses activation function transformations to control signal propagation, achieving training speed parity between vanilla and residual networks on ImageNet assuming the use of strong 2nd-order optimizers like K-FAC (Martens & Grosse, 2015). Zhang et al. (2022) extended ideas from DKS to a larger class of activation functions, achieving near parity in terms of generalisation as well.
The key quantity that is analysed in signal propagation is the DNN’s initialisation-time kernel, or more precisely, the approximate kernel given by the infinite width limit (Neal, 2012; Matthews et al., 2018; Lee et al., 2018; Yang, 2019). For MLPs, and for CNNs that use a Delta-initialisation (Balduzzi et al., 2017; Xiao et al., 2018), this kernel can be written as a simple recursion over layers that involves only 2D functions, facilitating a straightforward analysis.
Analogously to the case of MLPs, where signal propagation is judged by looking at the behavior of the (one-dimensional) kernel, signal propagation in transformers can be judged by looking at the evolution of these (high-dimensional) kernel matrices through the layers of the network. One situation we must avoid is where the diagonal entries rapidly grow or shrink with depth, which corresponds to uncontrolled activation norms and can lead to saturated losses or numerical issues. A more subtle form of signal degradation can occur where converges to a rank-1 matrix, which is known as rank collapse (Dong et al., 2021). Dong et al. (2021) showed that skip connections are essential to avoid the collapsed state: skipless transformers quickly converge to rank collapse at large depths, which we corroborate in Fig. 1 (top). Moreover, Noci et al. (2022) showed that rank collapse may lead to zero gradients for certain parameters in attention layers, hindering the trainablility of deep transformers. Thus, avoiding rank collapse is necessary for deep transformers to be trainable, and the question of whether one can train deep skipless transformers remains open.
In the present work we address this question, demonstrating for the first time that it is possible to successfully train deep transformers without skip connections or normalisation layers. To do so, we study the problem of signal propagation and rank collapse in deep skipless transformers, and derive three approaches to prevent it in Section 3. Our methods use combinations of: 1) parameter initialisations, 2) bias matrices, and 3) location-dependent rescaling, and highlight several intricacies specific to signal propagation in transformers, including the interaction with positional encoding and causal masking. In Section 4, we empirically demonstrate that our approaches result in trainable deep skipless transformers. On WikiText-103 and C4 datasets we show that using our main approach, Exponential Signal Preserving Attention (E-SPA), it is possible to match the training loss of standard transformers with our skipless ones by training for around 5 times longer. Moreover, by combining this approach with skip connections, we show that transformers without normalisation layers are able to match the training speed of standard ones.
Problem setting
where the shortcut and residual weights are typically both . In this work we will focus on skipless transformers with and , and on vanilla transformers, which are skipless transformers without normalisation layers. For simplicity we will also devote much of our analysis to attention-only models, which are transformers without MLP blocks (so that ).
Note that we restrict our analysis to decoder-only transformers in this work, as they have a simpler structure which is easier to analyse, and are widely used in practice. Also note that Eq. 1 corresponds to the “Pre-LN” (Baevski & Auli, 2018; Child et al., 2019) rather than the original “Post-LN” transformer (Wang et al., 2019). In Post-LN transformers, the normalisation operation is applied at the output of each MLP and attention block instead of at the beginning of each residual branch.
where is a large positive constant that zeros the attention coefficients corresponding to future tokens, making a lower triangular matrix.
As discussed in Section 1, Dong et al. (2021) showed that deep skipless transformers suffer from rank collapse, where the kernel matrix converges in depth to have rank 1, and Noci et al. (2022) showed that rank collapse can prevent trainability.
Moreover, Noci et al. (2022) demonstrated that rank collapse in transformers can occur in the absence of normalisation layers even with skip connections, and that downscaling the residual branch by setting can alleviate this issue. This latter observation is in line with previous findings concerning the benefits of downscaling residual weights in ResNets (Hanin & Rolnick, 2018; Zhang et al., 2018; Arpit et al., 2019; Hayou et al., 2021; Bachlechner et al., 2021) and transformers (Zhang et al., 2019; Xu et al., 2020; Huang et al., 2020; Touvron et al., 2021; Wang et al., 2022). Davis et al. (2021) showed that concatenation acts similarly to a downweighted skip as an alternative way to connect skip and residual branches. De & Smith (2020) noted that the interaction of standard skip connections and normalisations can also effectively downweight the residual branch to give better signal propagation properties, but only if the normalisation layer is placed on the residual branch, like in Pre-LN transformers. Such an effect does not occur for Post-LN transformers, where the normalisation layer is after the residual branch, and we indeed observe in Fig. 7 that Post-LN attention-only transformers also suffer from rank collapse at large depths. This may explain some of the training instabilities of Post-LN transformers that have been observed in practice (Xiong et al., 2020; Liu et al., 2020).
Constructing trainable deep transformers without shortcuts
To date, the only strategy for rectifying rank collapse in transformers relies on skip/shortcut connections, which “skip” around the trainability issues intrinsic to self-attention layers. We seek instead to tackle this issue directly. To do so, we first develop a better understanding of signal propagation through attention layers, then derive modifications from our insights to achieve faithful signal propagation in deep transformers, allowing them to be trained regardless of the use of skip connections.
To start with, we consider the simplified setting of a deep attention-only vanilla transformer, and suppose we are in a single-head setting () or a multi-head setting where the attention matrix does not vary across heads. If block has attention matrix at initialisation, then the final block’s representation takes the following form:
From this simplified formula for kernel matrices in deep attention-only transformers, we identify three requirements on :
must be well-behaved at each block, avoiding degenerate situations such as rank collapse and exploding/vanishing diagonal values.
must be elementwise non-negative (recalling that is constructed through the softmax operation Eq. 4).
should be lower triangular , for compatibility with causal masked attention.We describe compatibility of our methods with non-causal attention in Appendix A.
In Sections 3.1 and 3.2, we focus on finding attention matrices that satisfy our desiderata above, and demonstrate how to modify softmax attention to achieve these attention matrices in Section 3.3.
An obvious solution to our above requirements on is the trivial one: , where each sequence location attends only to itself. In this case, perfectly preserves the input kernel matrix, and will avoid rank collapse assuming is non-degenerate. Unfortunately, identity attention isn’t compatible with a viable solution to obtain trainable vanilla transformers. This is because to achieve we would need to saturate the softmax operation (Eq. 2) so that gradients do not pass to the query and key parameters, and the attention matrix stays close to identity during training.
To provide a partial solution to this that achieves identity attention matrix at initialisation yet is still trainable, we introduce our first approach, Value-SkipInit, based on SkipInit (De & Smith, 2020), and the related ReZero method (Bachlechner et al., 2021). In Value-SkipInit, we modify the attention operation , to
with trainable parameters and that are initialised to and respectively. Thus, at initialisation, the attention matrix is the identity. For transformers with MLPs blocks, this yields identical behaviour to a standard MLP acting on each sequence location independently at initialisation, so we can apply the DKS or TAT frameworks (Martens et al., 2021; Zhang et al., 2022) to achieve well-behaved signal propagation in the entire model.This is similar in principle to the Delta-initialisation for CNNs (Balduzzi et al., 2017; Xiao et al., 2018), and modifications to GNNs to enable compatibility with TAT (Zaidi et al., 2022). We note that Value-Skipinit is an approach to remove the standard skip connections in transformer blocks, Eq. 1, but can be construed as adding a skip connection around the value computation, and in that respect isn’t strictly “skipless”. Moreover, there is useful information contained in the positions of sequence locations that we do not employ when all attention matrices are identity at initialisation. As a result, we treat Value-Skipinit as a baseline for our main two methods, which we describe next.
2 Signal Preserving Attention methods
Returning to our requirements on in Section 3, we see that controlling their product is key to achieving faithful signal propagation. Given that we are now interested in non-identity , this becomes more difficult. To overcome this, we consider of the form
such that individual ’s cancel in the product, giving . Then if satisfies , we thus have
Assuming input embeddings are initialised independently, and no repeated tokens in the input sequence, we have in the large width limit,Because the embedding matrix is initialised with variance we rescale the embeddings it produces by to get a with ones on the diagonal. and can thus take . In practice, we will apply slight modifications to our methods to account for repeated tokens, as detailed in Appendix B.
Now if is lower triangular, then from Eq. 9 we see it is simply a Cholesky factor for the kernel matrix . By the uniqueness of Cholesky factors for PSD matrices up to sign, this means we just need to choose to be a family of well-behaved kernel matrices that satisfyConstraint (iii) is satisfied as lower triangular matrices are closed under multiplication and inversion. our non-negativity constraint (ii) on . We identify two such families, which give rise to our main method, Signal Preserving Attention (SPA):
Uniform (U-SPA): . Here, the kernel matrix has diagonal entries equal to 1 and off-diagonal entries equal to . The condition is required for elementwise non-negativity of the ’s, as shown in Theorem 1. Setting yields identity input kernel matrix, and as long as , we avoid rank collapse in the skipless attention-only setting.
Exponential (E-SPA): \big{(}\Sigma_{l}(\gamma_{l})\big{)}_{i,j}=\text{exp}(-\gamma_{l}|i-j|). Here, diagonal entries are again 1, but now off-diagonals decay exponentially for more distant locations, with decay rate . Thus, unlike U-SPA, E-SPA captures the notion of positional encoding, as the vector representations for nearby locations have larger inner products (i.e. are more similar to each other). The condition is required for elementwise non-negativity of the ’s, as established in Theorem 2. Setting yields the identity input kernel matrix, and rank collapse will be prevented as long as .
We state and prove Theorems 1 and 2 in Appendix I, and in particular provide a closed-form solution for in the case of E-SPA in Theorem 3, enabling cheap computation. From these theorems, we see that our proposed SPA approaches are viable as long as the kernel matrix values become progressively larger (in an elementwise fashion) as depth increases, as dictated by and . We find and to be good default choices for the final block, and describe how to vary and across depth in Appendix C. In Alg. 1 in Appendix D, we summarise how to construct a trainable self-attention layer with our main method E-SPA using ideas from Section 3.3, given an input-output decay rate pair (in the notation of Alg. 1). Given a decreasing set of decay rates , at block we set and .
In Fig. 1, we verify that our two proposed SPA schemes, U-SPA and E-SPA, successfully avoid rank collapse in attention-only vanilla transformers even at large depths. Moreover, because has diagonals equal to 1 for all , there is an implicit mechanism in these two schemes to control the representation vector norms across all sequence locations at deep layers. This means that they can be used with or without normalisation layers, as we will verify empirically in Section 4. Moreover, in Fig. 1 we see that E-SPA observes a recency bias as expected, where representation vectors for nearby locations have larger cosine similarity, as seen with positional encoding schemes like ALiBi (Press et al., 2022) (Fig. 7). As a result, even though all three of our approaches successfully avoid rank collapse, we expect E-SPA to outperform both U-SPA and Value-SkipInit.
3 Reverse engineering self-attention layers at initialisation
In Section 3.2 we identified two families of lower-triangular non-negative attention matrices which enable us to obtain well-behaved kernel matrices at large depth , and hence faithful signal propagation. It remains to show how we can actually realise these attention matrices via parameter initialisation and minor modifications of the attention mechanism.
which reduces to (with a data-independent attention matrix) at initialisation if the query-key dot product . We note that zero query-key dot products occur if one uses a scaling rather in the infinite width limit for independently initialised (Yang, 2019). In practice, to achieve zero initial query-key dot product, we can initialise either: 1) , 2) , or 3) both and to have small initial scale (which achieves approximately zero dot product). In our experiments we found these three options to all perform similarly, and decided to use option 1): . An empirical evaluation of the sensitivity to the choice of initial scale in option 3) is provided in Fig. 8.
In Eq. 10, acts as a non-trainable location-dependent rescaling, and is needed to realise arbitrary attention matrices since the softmax output is constrained to have row sums equal to 1.In SPA, also corrects for the fact that masked softmax attention tends to reduce the norms of representation vectors for locations at the end of the sequence. To see this, if we have kernel matrix with , and softmax attention matrix with row , then . But sums to 1 (as a softmax output), so with equality if and only if has exactly one non-zero entry (equal to 1). This holds for the first token (which can only attend to itself in masked attention) so , but not in general for later tokens, so will usually be less than 1, for . Given target kernel matrices with 1’s on the diagonal (which we have in SPA) its inclusion is akin to using RMSNorm at initialisation. However, this will gradually stop being true during training, leading to a (slightly) different model class. The additive pre-softmax biases will also have an effect on the model class, but this will be similar to that of the ALiBi positional encoder (Press et al., 2022), which too involves a non-trainable bias matrix being added to the logits. In our early experiments we tried including a trainable gain parameter on the bias matrix, initialised to 1, in order to better preserve the model class, but found that this didn’t have a significant impact on training performance.
We also note that this approach to controlling the attention matrix by zeroing the query-key dot product at initialisation is compatible with popular positional encodings like Relative (Shaw et al., 2018; Huang et al., 2019; Dai et al., 2019) and RoPE (Su et al., 2021), as discussed in Appendix E. Unless stated otherwise, we use RoPE in our experiments, and provide an ablation in Fig. 10.
4 Addressing MLP blocks and skip connections
Up until this point we have restricted our focus to attention-only skipless transformers, for the sake of simplicity. However, MLP blocks are an important part of the transformer architecture that we must also address. And we would like for our approaches to be compatible with skip connections too, both for the sake of generality, and because they might combine profitably.
Because MLP blocks operate on the representation vectors independently across locations, their effect on the kernel matrix can be easily computed using the standard limiting kernel formulas for MLPs (Neal, 2012) which are exact in the infinite width limit. In particular, there is a known that maps the kernel matrix of an MLP block’s input sequence to the kernel matrix of its output sequence. In principle, one could modify SPA to account for this change to the kernel matrix by taking , where and are the Cholesky factors of and respectively. Unfortunately, there is no guarantee that the resulting would be elementwise non-negative in general. In our experiments we thus elected to ignore the effect on the kernel matrices of the MLP blocks, approximating as the identity function, which we found to work well enough in practice. As discussed in Martens et al. (2021), this approximation becomes better as the MLP blocks tend to linear functions (for which is exactly identity), which will happen as we increase the shortcut weight relative to the residual weight (see Eq. 1), or when the MLP’s activation functions are close to linear. Notably, the latter condition will tend to be true when using DKS or TAT to transform the activation functions in deep networks. Note, when using DKS, decreases the size of off-diagonal elements of the kernel matrix, which intuitively will help to combat rank collapse.
When using skip connections as in Eq. 1, the output kernel matrix of a block is given by . One of the main ways that kernel matrices degenerate in DNNs is when their diagonal entries either explode or shrink with depth. As shown by Martens et al. (2021), this can be prevented in fully-connected and convolutional DNN architectures by rescaling the output of the activation functions (which happens automatically as part of DKS or TAT), and by replacing weighted sums with “normalised sums”, which yields normalised skip connections defined by the condition . When using normalised skip connections with U-SPA, both and will have the form for some , and thus so will . Moreover, we will have , which is less than . This means we can easily adjust U-SPA to be compatible with normalised skip connections by replacing in the formula with the Cholesky factor of . In the case of E-SPA, we note that won’t be of the form \big{(}\Sigma(\gamma)\big{)}_{i,j}=\text{exp}(-\gamma|i-j|) for some , even when both and are. To work around this, we use an approximation described in Appendix F, which seeks to make the combined effect on the kernel matrix of the normalised skip and attention block approximately equal to the effect of a skipless attention block.
Experiments
We now assess the capabilities of our proposed methods in training deep skipless and/or normaliser-free transformers. Our main experiment setting uses a transformer with 36 transformer blocks, which is deep enough that the effects of poor signal propagation and rank collapse render skipless training without modifications impossible. We begin our investigation on WikiText-103 (Merity et al., 2017) focusing primarily on training performance of our various methods, before moving to the larger C4 dataset (Raffel et al., 2019), where overfitting is not an issue. We use Adam optimiser (Kingma & Ba, 2014), as is standard for transformers, and train without dropout (Srivastava et al., 2014). Additional results and experimental details are provided in Appendices G and H respectively.
To start with, we verify that a standard deep transformer without skip connections is untrainable even with normalisation layers (LN) and transformed activations, and that our approaches remedy this. Fig. 2 compares vanilla transformers using our proposed SPA methods and Value-Skipinit to standard transformers both with and without skips, on a 36 block transformer. We clearly see that removing skip connections from a standard transformer makes it untrainable, with training loss plateauing around 7.5. This holds true too even when we use DKS to transform the GeLU MLPs, highlighting that the issue lies with the attention layers, which suffer from rank collapse as shown in Fig. 1.
On the other hand, all three of our approaches train even for vanilla deep transformers, with our E-SPA method outperforming U-SPA and Value-Skipinit. However, the default transformer with skips and LN still retains a training speed advantage compared to our skipless methods, mirroring the situation for CNNs without powerful second-order optimisers (Martens et al., 2021; Zhang et al., 2022).
In Table 1, we assess the effect of different activation functions in the MLP blocks, as well as the use of LN, in skipless transformers using our proposed methods. We see that at depth 36 we achieve good training performance for a range of activations: DKS-transformed GeLU, TAT-transformed Leaky ReLU, as well untransformed GeLU (Hendrycks & Gimpel, 2016), but not untransformed Sigmoid. We also see that layer normalisation is relatively unimportant for training speed, and can even be harmful with transformed activations when using SPA, which already has an inbuilt mechanism to control activation norms (as discussed at the end of Section 3.2).
In Fig. 3, we see that one way to match the training loss of the default transformer, without more iterations, is by using normalised skip connections. While this is perhaps unsurprising, we observe that our E-SPA method (left) matches standard training both with and without normalisation, whereas a standard Transformer with normalised skip connections (right) requires normalisation in order to match the training speed of the default Pre-LN.
So far, we tested our proposed methods on WikiText-103 (Merity et al., 2017), on which we observed overfitting without the use of extra regularisation. Therefore, we further compare our methods to standard transformers on a larger dataset, C4 (Raffel et al., 2019), where overfitting isn’t an issue. Importantly, we see similar trends across validationArgued by Nakkiran et al. (2021), the validation curves measure the “training speed” of the online setting. (Fig. 4(a)), training (Fig. 13) and downstream task (Table 6) performance, and so the benefits of our methods do extend beyond training. Due to the memory overhead of longer sequences, we use a 32-block transformer. One can notice that E-SPA performs the best among skipless transformers on all settings: training, validation and downstream tasks.
Moreover, in Table 2 we find that E-SPA with normalised skips and LN outperforms the default Pre-LN transformer, achieving 24.0 (vs 24.7) validation perplexity on C4 after 50K steps. Even without LN, E-SPA exactly matches the Pre-LN transformer in training speed, and outperforms a range of baselines designed to remove LN, including Stable ResNet (Hayou et al., 2021; Noci et al., 2022) and SkipInit (De & Smith, 2020), both with & without LN.
In Figs. 2 and 4(a), we observed that while our methods are able to train deep skipless transformers (a result which is unprecedented in the literature), there is a significant gap in training speed compared to a standard Pre-LN transformer (with skips). Martens et al. (2021) and Zhang et al. (2022) observed a similar training speed gap for skipless CNNs, and showed that such a gap can be closed by using more sophisticated second order optimisers like K-FAC (Martens & Grosse, 2015) or Shampoo (Gupta et al., 2018; Anil et al., 2020). As second order optimisers for transformers are not well established, we instead demonstrate that the training loss gap can be closed by simply training for longer in Fig. 4(b). We observe that our E-SPA method matches the training loss of a standard pre-LN transformer on C4 if one trains for around 5 times longer with Adam, in line with the findings from the convolutional case (Zhang et al., 2022). An equivalent plot for WikiText-103 is provided in Fig. 14.
Finally, as unmodified networks suffer from worse signal propagation properties at larger depths, it is natural to ask how our modified vanilla transformers perform as depth increases. In Table 3, we compare the training performance of different depths (36, 72, and 108) at a range of training step budgets on WikiText-103. We find that whilst the shallower depth 36 network trains fastest initially over 100K steps, it is matched at 400K steps and then surpassed at 1000K steps by the larger capacity depth 72 network. Moreover, our depth 108 vanilla transformer is able to close the gap in performance to its depth 36 counterpart with longer training, going from 0.3 at 100K steps to just 0.01 at 1000K steps.
Conclusion
We have shown for the first time that it is possible to successfully train deep transformers without skip connections or normalisation layers. To do so, we have proposed 3 approaches: E-SPA, U-SPA and Value-Skipinit, each of which control the attention matrices of a transformer to enable faithful signal propagation even at large depths. Our best approach, E-SPA enables deep vanilla transformers to match their standard counterparts with around 5 times more iterations, and also deep transformers without normalisation to match the training speed of standard ones. We hope that our work may potentially pave the way to new and improved architectures, and more research into improving the capabilities of deep learning in practice using insights from theory.
Reproducibility Statement
Pseudocode for our main approach, E-SPA, can be found in Alg. 1, using the notation and setup provided in Section 2. All experimental details can be found in Appendix H, including general and experiment-specific implementation details.
Acknowledgements
We thank Christos Kaplanis for helpful discussions during initial stages of this project, as well as the anonymous reviewers for their feedback. BH is supported by the EPSRC and MRC through the OxWaSP CDT programme (EP/L016710/1).
References
Appendix A Compatibility with non-causal attention
In Section 3, we focus on causal masked self-attention for two reasons. First, next-token prediction using causal masked self-attention is arguably the most popular setting for self-attention. And second, it is a more challenging setting to work with in terms of controlling signal propagation, due to the additional constraint that attention matrices must be lower triangular. In this section we describe how our methods can be made compatible with non-causal masked attention, where the attention matrices are no longer required to be lower triangular.
To start with, Value-SkipInit does not modify the softmax-attention computation and hence is already compatible with any form of attention. For our SPA methods, it is straightforward to extend to non-causal attention by changing in Eqs. 8 and 9 from being the Cholesky decomposition of to being the (symmetric) matrix square root of . In this case, for U-SPA it is possible to analytically calculate that in Eq. 8 will be element-wise non-negative if (exactly like the Cholesky case in Theorem 1). This is easy to see because the matrix square-root, inverses and products of uniform kernel matrices are all still uniform of the form (up to positive rescaling), so that will be too, and one simply needs to track and verify that it is positive. For E-SPA. we have verified empirically that the resulting will be non-negative if , just like the Cholesky case in Theorem 2.
Appendix B Modifications to SPA methods for repeated tokens
For simplicity, we assumed that our input sequences had no repeated tokens when presenting SPA in Section 3. This meant that we could take the input kernel matrix to be the identity, with zero off-diagonals, which was convenient for our construction of SPA. The effect of repeated tokens, e.g. if the word ‘cat’ occurs multiple times in the same sentence, is that our input kernel matrices will have non-zero off-diagonal entries, corresponding to entries where a token is repeated. This will impact the kernel matrices at deeper layers with SPA, particularly the diagonal values of , which we would like to able to control. In this section we discuss how we can modify our SPA approaches to account for the fact that we will often be working with sequences where a fraction of the tokens are repeated.
We stress that the general principle that all our methods (Value-SkipInit and SPA methods) follow is independent of the input kernel matrix (i.e. independent of duplicate input tokens): we seek to prevent the product of attention matrices from deviating away from the identity matrix and degenerating to a rank-1 matrix. From Eq. 6, we see that if the product of attention-matrices is rank-1 then regardless of the input kernel we will have a rank-1 output kernel i.e. rank collapse (Dong et al., 2021). On the other hand, if we control the deviation of the attention matrix product from the identity then no matter the input kernel, the output kernel will bear some similarity to the input kernel. So as long as the input kernel is non-degenerate and has full rank (regardless of duplicate tokens) so too will be the output kernel.
where is the product of attention matrices.
In SPA, we parameterise for corresponding to the Cholesky factors of some family of kernel matrices , with either uniform or exponentially decaying off diagonals:
U-SPA: for
E-SPA: \big{(}\Sigma_{l}(\gamma_{l})\big{)}_{i,j}=\text{exp}(-\gamma_{l}|i-j|)) for .
It is important that is constructed from two Cholesky matrices belonging to the same family, because Theorems 1 and 2 show in that case we will satisfy our non-negativity constraint on (which is computed through a softmax operation), and otherwise there is no prior reason to suppose that non-negativity will be satisfied.
Therefore, in an ideal world, would be an element of our family of kernel matrices, so that we can set to be the Cholesky factor of , satisfying .
Clearly, when , we have is a member of both the uniform and exponential families of kernel matrices, corresponding to uniform off-diagonals of or exponentially decaying off-diagonals with rate . In this case we can set too.
On the other hand, for , it is in general difficult to say much more about an individual sequence’s , given that different sequences will have repeated tokens in different locations. Moreover, the naive approach of ignoring the repeated tokens and treating leads to increasing diagonal values of (i.e. activation norms) at large depth without corrections, as shown in Fig. 6.We find that for sentencepiece tokenisations like we use in WikiText-103 and C4, but for character-level prediction like in EnWiki-8. This imbalance between blocks could be problematic for training dynamics, and also is incompatible with frameworks like DKS and TAT which suppose that diagonal values of are constant across blocks and locations, usually set to 1.
To circumvent this, we will derive our modifications by considering the average input kernel matrix (averaged over different sequences), under the assumption that repeated tokens occur independent of location:
By linearity of Eq. 11 in (and because we are controlling to be input-independent at initialisation in Section 3), it also follows that the average depth l kernel matrix satisfies:
so we can modify our SPA approaches to control the average kernel matrix instead.
For U-SPA, the situation is more straightforward as is a uniform kernel matrix with , so it suffices to simply let instead of .
For E-SPA, we are unable to view as having exponentially decaying off-diagonals, and so we keep which translates to and . Instead, to help us understand the effect of repeated tokens (to motivate our modifications), we first expand on Eq. 13 to simplify a little:
Note that all terms in Eq. 14 are easily computable (given that \big{(}\Sigma_{l}(\gamma_{l})\big{)}_{i,j}=\text{exp}(-\gamma_{l}|i-j|)) and has an analytic form provided in Lemma 1). So we can use Eq. 14 to compute the expected diagonal for each location (where the expectation is taken over different sequences), which we denote by a diagonal matrix :
Thus, we propose to replace with , setting by default. This means that our product , and Eq. 13 is updated to:
Though our modifications for repeated tokens only consider averages across different sequences, we find that for individual sequences they still lead to well behaved diagonal values of , as shown in Fig. 6. Moreover, we find that their effect on off-diagonals is still favourable for individual sequences, as shown in Fig. 1.
In this section we describe how we set the uniform off-diagonals and the exponential decay rates in SPA, at different depths.
Recall that Theorems 1 and 2 show that we are free to choose and such that increase with depth and decrease with depth. Moreover, (or the shared token fraction , as per Appendix B) and at the input layer, whilst we have found and to be good default values for the last block.
For U-SPA, in terms of setting how vary with depth, we tried different polynomial rates of increase with depth from to , but did not observe a noticeable difference in performance across different rates so chose to increase from to linearly in depth.
For E-SPA, we choose so that the diagonal elements of the attention matrices are constant across blocks, akin to using a constant shortcut weight over different blocks. From Theorem 3, we have that the diagonal entries of satisfy:
where with inverse . This means that for a set of positive decreasing , there exist a corresponding set of decreasing with values between and .
Because , we have , and likewise for a given we can compute .
Thus to have constant diagonal values of over different blocks , we choose to set
We found this scheme to work well empirically, and note the similarity to other works which have discussed the choice of how to scale branches in residual architectures (Hayou et al., 2021). We leave further study of choosing how to set and to future work.
Appendix D E-SPA algorithm
In Alg. 1 we present pseudocode to construct a trainable E-SPA masked attention layer.
Appendix E Compatibility of SPA with existing positional encodings
In Section 3.3, we showed how to control the attention matrix at initialisation, using bias matrices and location-dependent rescaling, as well as making the query-key dot product, , zero at initialisation. This scheme is used in our SPA methods. We detailed several ways to achieve this, and in practice we chose to initialise to zero, and initialise as usual (Gaussian fan-in or orthogonal).
In this section we show how zero initialising the query-key dot product is also possible when using two standard positional encodings: RoPE (Su et al., 2021) and Relative (Shaw et al., 2018; Huang et al., 2019; Dai et al., 2019). This means that we can use our methods in Section 3.3 in combination with these positional encoders.
Let us denote the unscaled query-key dot product as .
For RoPE, the query-key dot product, , is modified from
For Relative positional encoding, we take the scheme from (Dai et al., 2019). In that case, the query-key dot product, , is modified from
Appendix F Using normalised skip connections with E-SPA
In this section, we describe how to combine our E-SPA method with normalised skip connections in our attention blocks:
where is the skipless setting that our methods are originally designed for.
The general gist is that we will look at the combined effect, on the kernel matrix after the residual attention block, of the residual branch and the diagonal terms in the attention matrices for preserving signal propagation, and approximate the combination to match the setting where we are without skips, i.e. .
To combine E-SPA with normalised skip connections, we consider how the cosine-similarity between two locations , is affected through the normalised skip connection. Suppose we have an input kernel matrix:
where and is the width. Then after the residual attention block, Eq. 18, we now have:
Either we have to be an orthogonal matrix sampled uniformly at random from the Haar measure, or . In both cases we have going towards , and the cross terms Eq. 20 converging to 0 for large .
Thus, we can consider the large approximation:
Then, if we look at the inner product between the and locations for locations such that , we have:
where is the constant diagonal of the attention matrices in E-SPA, c.f. Theorem 3, and .
Moreover, we have defined as it is a term that is not possible to control with only knowledge of and will vary from sequence to sequence, and hence we argue can be discarded from a signal propagation perspective. For example, one could consider a sequence where the kernel matrix has all off-diagonals equal to 0 apart from , and could be sufficiently distant from such that there is no for which and are large at the same time (which can be seen using the analytic form of given in Theorem 3).
Thus, from Eq. 23, we see that after an residual attention block, an input cosine similarity of between locations is diluted by a factor of
with shortcut weight and attention diagonal probability . Our approximation then seeks to preserve this factor when to match the skipless case.
Now, in the skipless case described in Section 3, at block we suppose that the incoming kernel matrix has exponentially decaying off-diagonals with rate , and that we construct the attention matrix so that the output kernel matrix has exponentially decaying off diagonals with rate . From Theorem 3, we see that this gives diagonal entries of to be:
where .
Thus, to preserve to match the case for , if we have shortcut weight we need the attention matrix diagonal probability to satisfy:
In turn, when we have shortcut weight , this means that we need to choose our outgoing decay rate at block , , such that:
Inverting the definition of , we see we need to set :
To summarise, if we have a sequence of decreasing exponential decay rates, our proposed approximation when using shortcut weights at block is, using the notation of Alg. 1, to set as normal, and to set from Eq. 25 in order to preserve the signal propagation from the combined residual attention block Eq. 23. We see that this approximation reduces to the standard skipless setting when , as a sanity check.
Appendix G Additional results
In this section we present additional results and ablations that were not included in Section 4.
In Fig. 7 we plot the evolution of normalised kernel matrices for transformers with skips and or RMSNorm normalisation, in addition to those for vanilla transformers in Fig. 1. We see that both skipless with RMSNorm (fourth row) and Post-LN (bottom row) converge to rank collapse at larger depths. The degeneration of skipless with RMSNorm is expected from the results of (Dong et al., 2021), and while the convergence to rank collapse is slower for Post-LN, it is still expected (Hayou et al., 2021; Noci et al., 2022). This is because the residual and shortcut branches are effectively given a constant weighting at all blocks in Post-LN, even as the network’s depth increases.
On the other hand, Pre-LN observes sensible signal propagation even at depth 100, as the positioning of the LN in the residual branch effectively downweights the residual branch at later blocks (De & Smith, 2020). Likewise, Pre-LN with skip weight also observes faithful signal propagation, because each block is effectively downweighted. This effect means that standard Pre-LN’s (fifth row) kernel matrix increases elementwise faster with depth than Pre-LN with normalised skips (sixth row), despite both kernel matrices being qualitatively similar at block 100. Note, all methods besides our SPA methods used ALiBi positional encoder (Press et al., 2022) (detailed in Section H.2) and we observe both Pre-LN kernel matrices also have a recency bias, like our main method, E-SPA (third row).
In Table 1 we compared the training speeds for our skipless methods, and found that E-SPA outperforms both U-SPA and Value-Skipinit. In particular, we found E-SPA with a DKS-transformed GeLU activation without LN to perform best. In Table 4, we present the corresponding results but for validation perplexity. Again, we see that E-SPA is the best performing of our attention modifications, but in this case TAT with Leaky ReLU and no LN matches or outperforms DKS with GeLU.
Recall in Section 3.3 that for attention layers using SPA we seek to initialise weights such that the query-key dot product, , is zero or small at initialisation. In our main experiments, we achieve this by initialising , and letting to be initialised as normal. In Fig. 8, we assess the sensitivity of our E-SPA scheme to the scale of non-zero attention dot product, when are both orthogonally initialised but with scale (i.e. at each layer, both are initialised as two independent uniform orthgonal matrices multiplied by ). We see that for small initialisation scales, there is little effect of varying initial scale, but that training performance degrades at larger scales, when our attention matrix reverse engineering in Section 3.3 (which expects small or zero query-key dot product at initialisation) is less precise.
Recall in Section 3 that our kernel matrix evolution for attention-only transformers is exact at finite widths using orthogonally initialised weight matrices, and will be approximate using standard (Gaussian) fan-in initialised In Fig. 9, we ablate over using orthogonally initialised weight matrices, compared to Gaussian fan-in initialisation. Across activations, we see that E-SPA with orthogonal initialisation slightly outperforms Gaussian fan-in initialisation (by around 0.15 train loss).
As noted in Appendix E, our methods are compatible with several standard positional encodings, and by default all our experiments use the popular RoPE (Su et al., 2021) positional encoding. In Fig. 10 and Table 5, we assess the effect of removing positional encodings, on training and validation performance respectively, in our skipless methods. We see that all methods are improved when combined with RoPE, however the improvement is most mild in E-SPA, which as discussed has an in-built recency bias akin to a positional encoder. Moreover, E-SPA without additional positional encoding still outperforms all other approaches, including U-SPA (which on its own has no notion of position) with RoPE.
In Fig. 11, we provide an equivalent plot to Fig. 3 using a Stable ResNet (Hayou et al., 2021) rescaling of the shortcut weights (). In this case, the shortcut weight is always and the residual weight is uniform across blocks and scales as in depth . Hayou et al. (2021) showed that such a scaling leads to non-degenerate signal propagation in large depth MLPs/CNNs without normalisation, and Noci et al. (2022) showed that the stable scaling prevents rank collapse in transformers without normalisation. We see in Fig. 11 that for large enough , the stable residual weighting with normalisation matches the training speed of the default transformer (which is unsurprising given that once it is exactly default Pre-LN). However, there is a small but consistent gap without normalisation (with optimal ). Here, , so , or alternatively if we count the 72 nonlinear layers (one self-attention and element-wise nonlinearity for each transformer block), we have .
Recall that in a transformer block Eq. 1, there are two distinct skip connections: one for the self-attention block and one for the MLP block. Moreover, we observed a training speed gap when we remove both skip connections in Fig. 2. This leads us to ask if it is possible that only one of the skips, MLP or attention, is causing this gap in training speed. In Fig. 12, we investigate this by varying the MLP shortcut weight for skipless attention blocks (left) and varying the attention shortcut weight for skipless MLP blocks (right). For all skip connections (for both MLP and self-attention blocks), we use a normalised skip connection (. We observe that removing either skip connection results in a comparable loss of training speed (the default Pre-LN on WikiText-103 obtained train loss of 1.76 after 100K steps), although having one is still better than having neither. Moreover, we observe that the attention and MLP blocks prefer slightly different shortcut weights, with dense shortcuts performing better on slightly lower weightings ( or ) compared to attention shortcuts ( or ), whereas our experiments in Figs. 3 and 11 use a joint weighting for both.
In Fig. 13 we compare the training performance of our various vanilla transfomers to the default Pre-LN transformer on C4. This is akin to Fig. 4(a), which compared validation performance.
Typically, transformers are pre-trained on a large corpus of data before evaluation on a set of downstream tasks. To assess whether our conclusions about training performance in pre-training transfer over to downstream tasks, we assess models trained on C4 on 5 common sense downstream tasks: BoolQ (Clark et al., 2019), HellaSwag (Zellers et al., 2019), Winogrande (Sakaguchi et al., 2020), PIQA (Bisk et al., 2020), and SIQA (Sap et al., 2019). These datasets are commonly used to evaluate large pre-trained transformers (Brown et al., 2020; Rae et al., 2021; Smith et al., 2022; Hoffmann et al., 2022).
In Table 6, we see that the conclusions of pre-training on C4 largely carry over to the downstream tasks:
Among transformers trained for the same number of steps (50k), E-SPA beats U-SPA and Value-SkipInit each on 4 out of 5 downstream tasks. However, the default transformer outperforms skipless transformers on all tasks with the same amount of training.
With around 5 times longer training (200K and 300K steps), E-SPA achieves similar performance on downstream tasks to the standard transformer, outperforming on 2 out of 5 tasks (Winogrande and PIQA).
In Fig. 14 we see that 4.5x more training allows a vanilla E-SPA transformer to match the validation performance of a standard transformer on WikiText-103. This mirrors our findings on C4 in Fig. 4(b).
Fig. 15 shows the evolution of the empirically-computed normalized kernel matrix during training of a (finite width) vanilla E-SPA transformer on WikiText-103. The network has depicted has both attention and MLP blocks, which use DKS-transformer GeLU activations. We note that the untrained network shows good agreement with Fig. 1, despite the fact that Fig. 1 is computed for an attention-only network in the infinite width limit. We also note that while significant changes to the kernel matrix occur during training, it retains the property of being larger close the diagonal.
Appendix H Implementation details
We first describe all additional general implementation details for our experiments, before going into details relevant for individual results.
We present experiments on WikiText-103 (Merity et al., 2017) and C4 (Raffel et al., 2019). For both datasets, we use the SentencePiece tokeniser (Kudo & Richardson, 2018) with vocabulary size . For WikiText-103 we use sequence length of 512 for both training and validation. For C4, the sequence length is 2048 for both.
In all our experiments apart from Table 3, the model width across all blocks. We use 8 head multi-head attention, so that . The MLP block consists of a single hidden layer with width , with input and output dimensions both equal to . All our experiments use the Pre-LN transformer block Eq. 1, rather than Post-LN. On WikiText-103, we use a 36 block transformer by default for all experiments apart from Table 3. On C4, we use a 32 block transformer due to memory constraints. Any normalisation layer we consider is RMSNorm (Zhang & Sennrich, 2019), which is simpler than Layer Normalisation (Ba et al., 2016) and is commonly used in transformers (Rae et al., 2021). By default, our models use RoPE positional encoder (Su et al., 2021) (apart from the ablation in Fig. 10).
By default, all weight matrices are fan-in initialised with . The two exceptions for this are: 1) when using orthogonal intiailisation, we use the scaled-corrected uniform orthogonal initialisation (Martens et al., 2021) with scale (which for square matrices is just an orthogonal matrix sampled from the Haar measure Meckes (2019) multiplied by ), and 2) for the parameter matrix immediately after the activation we set to take input activation norm (“q-values” or diagonals of Gram matrices) of 1 to output activation norm of 1. In the latter case, by construction for activations transformed by DKS/TAT Martens et al. (2021); Zhang et al. (2022). All bias parameters are initialised to 0.
For DKS (Martens et al., 2021) we set slope parameter , and for TAT (Zhang et al., 2022) with leaky ReLU, we set . Both values were chosen by a small hyperparameter sweep on WikiText-103. The DKS and TAT transformations are chosen without consideration of the attention blocks, where the transformer can be viewed as an MLP (potentially with residual connections). Unless stated otherwise, all skipless transformers used a DKS-transformed GeLU as the nonlinearity in the MLP block by default.
We use Adam optimiser (Kingma & Ba, 2014) with global gradient clipping of 0.1 by default (Pascanu et al., 2013). We do not use weight decay in our experiments.
We use mini-batching of 16 sequences for WikiText-103 and 8 for C4, due to memory constraints. Unless stated otherwise, we train for 100K steps on WikiText-103 and 50K steps for C4.
H.2 Additional implementation details for individual experiments
We calculate the kernel matrix evolution directly in kernel matrix-space, where . Our input kernel matrix is sampled assuming a fraction of repeated tokens, with value of if the token is repeated, and 0 else. For all configurations of skip/normalisation/attention modifications corresponding to a row of Fig. 7, we use 8 heads.
We now detail how each operation in any configuration of Fig. 7 affects the kernel matrix. From this it should be possible to reconstruct the kernel evolution for any row in Fig. 7. The 3 possible operations are: 1) attention, 2) skip connection, or 3) LN/RMSNorm operation.
Attention Because our SPA methods (second and third rows) are agnostic to the number of heads at initialisation (all attention matrices in a self-attention block are the same across heads at initialisation), we apply Eq. 6 directly, so that a single attention block amounts to for attention matrix and incoming kernel matrix . Our E-SPA method uses and our U-SPA method uses .
For all other rows, the self-attention operation uses ALiBi (Press et al., 2022), a popular positional encoder which uses head-dependent pre-softmax bias matrices. More specifically, from the default pre-softmax bias matrices given by ALiBi (for 8 heads), we obtain 8 attention matrices using the softmax operation (which is exact assuming zero query-key dot product at initialisation). Because the different heads in transformers are typically concatenated along 8 equal size fractions of the total width , the kernel evolution of an attention block on kernel matrix with 8-head ALiBi corresponds to (Martens et al., 2021):
For a skip connection with shortcut weight and residual weight , if denotes the output of a kernel matrix after an self-attention operation (i.e. from point 1. above), then an incoming kernel matrix gets mapped to:
For a normalisation operation, the incoming kernel matrix gets mapped to:
For experiments on WikiText-103 with 100K steps, for our U-SPA transformers we tuned , and for our E-SPA transformers we tuned . For all other settings (longer/deeper training on WikiText-103 or any C4 experiment), we used the default and . All hyperparameters throughout our work are tuned based on training loss.
All experiments with skip connections (i.e. shortcut weight for either the MLP or attention block) use untransformed GeLU activation in the MLP blocks. We combine E-SPA with normalised skips as described in Appendix F. We note that for high shortcut weights and small final decay rate , then the value of attention matrix diagonal , Eq. 24, may not be real. This is because the input to the square root in Eq. 24 may be negative. To get by this, we tune when using a 36 block transformer on WikiText-103 and for Table 2, which used a 32 block transformer on C4.
For Table 2, the normalised skip connections had separately tuned attention and MLP shortcut weights, as we observed a difference in the optimal shortcut weight for self-attention vs MLP blocks in Fig. 12. We tuned the attention shortcut weight in the range and the MLP shortcut weight in the range . Likewise, both the stable residual weights were tuned in separately for self-attention and MLP skips. The selected shortcut/residual weights (using validation performance) are presented in Table 7.
Due to memory constraints, for our deeper networks in Table 3 we use width rather than , with 8 heads to give . All depth scaling runs used a DKS-transformed GeLU with .
Appendix I Theoretical results
In this section we state and prove our theoretical results, including Theorems 1 and 2:
(Non-negativity for U-SPA) Let and , with respective (positive) Cholesky factors and . Then if , we have is elementwise non-negative.
(Non-negativity for E-SPA) Let matrices and with respective (positive) Cholesky factors and . Then if , we have is elementwise non-negative.
We prove Theorem 1 second as it is more involved.
We actually prove Theorem 2 as a corollary of Theorem 3, which provides the analytic form for .
Let matrices and with respective (positive) Cholesky factors and . Then if , we have takes the following form:
where , and likewise
From Theorem 3, we have the analytic form of . Clearly are positive, and moreover if , then is non-negative.
Finally, because when , then is non-negative too. ∎
To prove Theorem 3, we first compute what the analytic form of Cholesky factor for takes in Lemma 1.
Let with (positive) Cholesky factors such that . Then, we have:
It is clear that is positive semi definite, as it is the covariance matrix of a stationary Ornstein-Uhlenbeck process, hence a Cholesky factor must exist. We now show that it is , Eq. 28.
If we define , then we have:
By the uniqueness of (positive) Cholesky factors, the proof is complete. ∎
We want to show . This is clearly true for the top diagonal .
We now show this for the rest of the first column, when :
Applying the geometric sum to Eq. 29 yields:
I.2 Proof of Theorem 1
Before diving into the actual proof of the theorem we will derive several useful properties and notations. First we make a slight notational change from the main text of the theorem by replacing , which depends on , with , where represents the size of the matrix, while replaces . In addition, since the case of (or in the rest of the proof) is trivial, since then the resulting matrix is the identity, we will restrict ourselves to dealing with the case where . In any of the mathematical derivations we will denote with capital English letters (e.g. ) any temporary expressions, that will be expanded on the following lines. Note that these are never general definitions, so they might be used multiple times for different expressions.
We will denote vectors and vector functions in bold and scalars and scalar function in standard font.
Let be the -dimensional vector with only ones:
We will denote with the -dimensional vector with only ’s:
First, we define the linear map as:
The vector is an eigenvector of with an eigenvalue .
Directly calculating the -th entry of the product gives:
. ∎
The vector is an eigenvector of with an eigenvalue .
.
.
Further we define the following useful functions:
I.2.2 Lemmas
If then all entries of the vector are non-negative - .
If then the function is non-negative.
If then the function is negative.
The partial sum the functions from to will be denote by , which from Lemma 4 follow are always negative.
I.2.3 Main proof
First for we have that , which implies that and the condition is trivially satisfied.
Now assuming that the statement is true for all integers up to , we will prove that it also holds for :
Using the fact that is a lower triangular and non-negative matrix by the inductive assumption combined with Lemma 2 and the fact that it follows that is also a lower triangular non-negative matrix.
I.2.4 Proof of Lemma 2
First we will inspect the evolution of as we increase :
Using the definition of from Lemma 3, Lemma 4 and Definition 2 we have that:
Thus proving the lemma reduces to proving that:
Expanding on the equation that we need to prove is positive:
Now we turn our attention to the sum in the middle:
Let’s define the partial sum in the brackets as:
First for , we have that which clearly satisfies the above equation. Assuming this is true for , we will now show it holds for :
Which concludes the proof. This now means that:
From the fact that it impliest that is a decreasing function of , hence to we only need to show that for any fixed and is positive.
For this, we need to take the limit of the second and third term in the above equation.
where the last line is true since the denominator is of higher degree in .
Taking the limit of the second term corresponds to computing the limit:
All of the limits must be taken for fixed and (e.g. we can’t have them approach or simultanously with ).
I.2.5 Proof of Lemma 3
From the definition of and the fact that and are non-negative functions we can conclude that:
Denoting with the left hand side of the second equation we have:
Hence for we have that .
I.2.6 Proof of Lemma 4
First we will prove that
With this we can now conclude that
Hence with this we can conclude that .