The Shaped Transformer: Attention Models in the Infinite Depth-and-Width Limit

Lorenzo Noci, Chuning Li, Mufan Bill Li, Bobby He, Thomas Hofmann, Chris Maddison, Daniel M. Roy

Introduction

Pre-trained large language models have experienced a remarkable increase in popularity due to their eerily human-like ability to puzzle through complex reasoning tasks, solve coding challenges, and produce pages of logically sound text . Arguably, the Transformer is the foundation of these successes . Recent research has found evidence for scaling laws, linking the performance of these architectures to their parameter counts and the quantity of training data, fueling the desire to train deeper and wider models on ever larger datasets in order to unlock new levels of performance .

Bundled with the increased expressivity of deep architectures, however, is increased numerical instability, both in the forward pass and gradients, which hinders training. One of the clearest examples of instability is the so-called rank collapse phenomenon – the observation that, in Softmax-based attention models, the network’s representation of different tokens tend to perfectly align at large depth. The resulting poorly conditioned covariance and correlation between tokens leads to exploding and/or vanishing gradients at initialization, disrupting gradient updates of the affected parameters. This situation violates a well-known guiding principle from the literature of deep signal propagation: a stable covariance is a necessary condition for stable training . In fact, the instability of Transformers is evident when considering the critical role of hyperparameter tuning and the judicious use of normalization layers. In this work, we study Transformers in a novel infinite limit, rectify sources of instability with a novel modification, and derive the SDEs characterizing the covariance and output distribution.

Scaling limits have been used successfully to provide guidance on architecture and tuning hyperparameters settings . Our work represents a contribution in this direction. The ability to use such limits to diagnose instabilities depends on their tractability and faithfulness to real-world (finite) networks. In this regard, not all limits are created equal. In particular, the faithfulness of scaling limits depends critically on how other parameters are scaled with width. One of the simplest (and thus most popular) limits to work with – the “NTK” limit – treats the depth of the network as fixed. As a result, at initialization, this limit does not accumulate sufficient random fluctuations over the depth of the network, leading to deterministic covariance matrices that do not agree with those of standard (finite) networks. Such networks have another defect: they are incapable of learning features in the limit . Various other limits have been studied, towards identifying tractable yet faithful models of initialization and/or training. These include mean field limits and the perturbative regime .

This work operates in a relatively new regime – the proportional infinite depth-and-width limit – where depth dd and width nn diverge as the ratio d/nd/n tends to a positive constant. This limit, first analyzed by , has been the recent subject of study in the context of neural network . A related line of work also studied the Lyapunov exponent for products of random matrices . This regime retains the network’s stochasticity and, at initialization, has been shown to closely resemble the behaviour of finite architectures, yet still yield a relatively simple limiting description, expressible in terms of stochastic differential equations . In this work, we fully characterize the initial output distributions of a network with skip connections and Softmax-based attention mechanisms, in the proportional infinite-depth-and-width limit.

Inspired by the idea of shaping activation functions , our theoretical approach finds an adequately modified attention mechanism via its SDE limit. Our modification involves making the attention matrix closer to the identity, and appropriately choosing the temperature parameter τ\tau, which re-scales the logits of the Softmax. Similar to shaping activation functions, the temperature scaling we devise linearizes and reduces the saturation of the Softmax, a known source of training instability in Transformers . In order to model the feedforward layer of a Transformer’s block, we extend existing results to derive an SDE for the proportional limit of shaped-ReLU feedforward multi-layer perceptrons (MLPs) with skip connections. Combined, we fully characterize the output distribution of a Transformer with shaped non-linearities (Corollary 4.3).

Notably, our modification successfully prevents a poorly conditioned covariance matrix, whereas the vanilla Softmax-based attention model without LayerNorm fails in this regard, and the corresponding Pre-LN architecture provides only marginal improvements (see Figure 1). Given that our modification is inspired by previous work on shaping activation functions, we coin the terms shaped attention for the proposed attention mechanism and shaped Transformer for the overall architecture that includes the MLP block and residual connections. Through simulations (e.g., Figure 1), we show that the limiting neural covariance SDE approximates the distribution of finite-size Transformers with shaped attention mechanism surprisingly well. We also provide preliminary training experiments for our proposed shaped attention architecture on standard language modeling tasks, demonstrating the feasibility of the new architecture in practice (see Section 5 and Appendix D).

In summary, our contributions are as follows:

We study the effect of skip connections in the proportional limit, showing that under a precise relation between the scaling parameters of the shortcut and residual branches, the feature covariance converges to the solution of a weighted version of the neural covariance SDE for MLPs (Theorem 3.2). The dependence on the depth-to-width ratio implies the existence of a stable non-commutative limit for residual networks, complementing the commutative limit studied in .

We propose shaped attention, where we modify the Softmax-based attention mechanism to be a perturbation of the identity. We demonstrate that shaped attention successfully prevents the degeneracy of correlation in contrast to existing Transformer architectures (Figure 1). The enhanced stability in the forward pass is reflected in the gradients, which are also stable with depth, as we empirically show in Figure 2.

For the proposed shaped attention architecture, we derive the neural covariance SDE characterizing the initial distribution in the proportional limit (Theorem 4.2). Consequently, we provide the first characterization of Transformer-type architectures, i.e. the shaped Transformer, in the large depth-and-width regime (Corollary 4.3).

We provide simulations to validate the theory and to interpret the effects of network hyperparamaters on the covariance matrix of the shaped Transformer. Specifically, we study finite time stability of the SDE and provide explicit guidance on hyperparameters to prevent numerical instability.

The paper is organized as follows: In Section 2, we provide the basic setup and some background on existing results. In Section 3, we generalize the SDE results of to include skip connections. This serves as a model to understand the effect of skip connections in isolation from the attention model. In Section 4, we present our main result, first pinpointing the origins of instability in the Softmax, then showing how the modifications underlying shaped attention allow us to derive a non-trivial SDE limit. Finally, in Section 5, we discuss the implications of our results and some future directions. Proofs of all theorems and additional experiments are deferred to the Appendix.

Background

Neural Covariance.

Stabilizing the Effect of Non-Linear Layers.

Central to the issue of degeneracy of the neural covariance are commonly used non-linear activation functions that severely deviate from the identity. The recent line of work of Deep Kernel Shaping (DKS) addresses the issue by considering the cumulative amount of non-linearity throughout layers, and shaping the activation function by making it closer to the identity map. Inspired by this line of work, devise an initialization for Transformers that avoid the rank collapse problem without the aid of skip connections or LayerNorm.

In an alternative approach, the line of work behind Stable ResNets considers scaling the residual branches by γ=1/depth\gamma=1/\sqrt{\text{depth}}, and postulates this scaling is sufficient to stabilize the neural covariance with minimal assumptions on the activation function. adopts this scaling to give precise formulas on the expected covariance of a Transformer at initialization. In this work, we consider γ\gamma constant in width and depth, and derive a complementary limiting result.

The Proportional Infinite-Depth-and-Width Limit.

where the formulae for coefficients bReLU,Σlinb_{\text{ReLU}},\Sigma_{\text{lin}} can be found in Theorem 3.2.

We note that the output neuron distributions are directly recovered as a conditional Gaussian with covariance VTV_{T} for T=dnT=\frac{d}{n}, in a similar spirit as the neural network Gaussian process (NNGP) results . For example, the ii-th output Xout,iX_{\text{out},i} conditioned on VdV_{d} are asymptotically iid. N(0,VT)\mathcal{N}(0,V_{T}) as d,n→∞d,n\to\infty. The reader is referred to Appendix A for more technical background on the covariance SDE and the convergence result.

While the existing results are limited to initialization, we remind the reader that this is a necessary step before we can study training dynamics. In particular, the NNGP techniques developed for infinite-width networks at initialization were directly used to study the training dynamics in the same limit . We will provide further discussions on this topic in Section 5.

Warm-Up: a Neural Covariance SDE for ResNets

To understand the effect of skip connections, it is helpful to look at a simplified model composed of a shaped ReLU-activated layer and skip connections:

We will next define the notion of convergence for our covariance matrices and state our first main result. We refer the reader to Appendix A for more precise details on the Skorohod topology.

where bres(V)=γ2bReLU(V)=γ2[ν(ραβ)VααVββ]α≤βb_{\text{res}}(V)=\gamma^{2}b_{\text{ReLU}}(V)=\gamma^{2}[\nu(\rho^{\alpha\beta})\sqrt{V^{\alpha\alpha}V^{\beta\beta}}]_{\alpha\leq\beta} with ραβ=Vαβ(VααVββ)−1/2\rho^{\alpha\beta}=V^{\alpha\beta}(V^{\alpha\alpha}V^{\beta\beta})^{-1/2} and

furthermore, Σres(V)=2γ2Σlin(V)=2γ2[VαδVβω+VαωVβδ]α≤β,δ≤ω\Sigma_{\text{res}}(V)=2\gamma^{2}\Sigma_{\text{lin}}(V)=2\gamma^{2}[V^{\alpha\delta}V^{\beta\omega}+V^{\alpha\omega}V^{\beta\delta}]_{\alpha\leq\beta,\delta\leq\omega}.

Notice how the limiting SDE closely resembles the MLP case (Equation 3), which is recovered exactly when γ=1\gamma=1. The only difference is the extra 22 factor, which comes from the fact that in our definition each layer has effectively two times the number of weight matrices than the standard formulation for MLPs. As the drift depends solely on the nonlinearity, and the diffusion depends soley on the random weights, only the diffusion variance is doubled. The residual branch parameter γ<1\gamma<1 dampens both the drift and the variance of the Brownian motion by γ2\gamma^{2}, thus it can be interpreted as a time change. In other words, the effect of γ\gamma at initialization is equivalent to reducing depth-to-width ratio, inline with existing intuitions that ResNets have a lower “effective-depth” . To visualize the stabilizing effect of γ\gamma on the distribution, in Figure 3 (right) we plot the 95th percentile correlation as a function of γ\gamma. The increasing trend indicates a larger probability of perfect alignment between two tokens. In Figure 3 (left) we plot the densities of both the residual SDE and the corresponding residual network for various values of γ\gamma. Notice how the samples from the SDE well-approximates the histogram of a finite network.

Neural Covariance SDE for Softmax-Based Attention

2 Shaped Attention

To shape the Softmax-attention mechanism as a perturbation of the identity matrix, we propose the following modifications which we call the shaped attention In principle, it could be possible to have a close-to-identity Softmax matrix when the logits are large. However, this regime also corresponds to a very saturated Softmax, thus making training unstable . As a result, we will avoid this direction in this work.

The shaped attention presents three modifications to the Softmax attention in Equation 2. Firstly, the zero-order term m−111⊤m^{-1}\mathbf{1}\mathbf{1}^{\top} of the Taylor expansion (Equation 8) is removed as it causes a non-infinitesimal drift in the Markov Chain that ultimately leads to instability in the covariance (see Section 4.1). Secondly, we also observe that when τ\tau is very large, the centered Softmax is a perturbation around zero. To recover an approximate Euler-update like in Equation 7, we simply add back the identity matrix. By biasing the attention matrix towards the identity, we encourage each token to self-attend. This type of modification was also previously considered by . Finally, the Softmax’s temperature is chosen to scale as τ=τ0nnk\tau=\tau_{0}\sqrt{nn_{k}}, for some constant τ0>0\tau_{0}>0, which guarantees a non-degenerate limit as (d,n)→∞(d,n)\to\infty (Theorem 4.2). Note that the extra n\sqrt{n} term is a departure from the standard parameterization.

In Figure 4, we show how removing any of the proposed changes individually alters the neural covariance structure, which becomes degenerate for large depths, while the proposed modifications remain stable. We stress that here for simplicity we focus on attention without masking. Shaped attention can be extended to include masking (e.g. casual masking) by centering each i-th row of the Softmax matrix by a different factor 1/mi1/m_{i}, where mim_{i} is the number of un-masked tokens in the i-th row.

3 Main Result – Neural Covariance SDEs for Shaped Attention Models and Shaped Transformers

Before we state our main results, we will first define a weakened notion of convergence, which is required whenever the drift and covariance coefficients are not Lipschitz. This was also required for the case of shaped MLPs with smooth activations .

We say the covariance V(n)V^{(n)} converges locally to VV if the stopped process {Vt∧Tr(n)}t≥0\{V^{(n)}_{t\wedge T_{r}}\}_{t\geq 0} converges to {Vt∧Tr}t≥0\{V_{t\wedge T_{r}}\}_{t\geq 0} in the sense of Definition 3.1 for all stopping times of the form Tr=inf⁡{t>0:∥Vt∥≥r}T_{r}=\inf\{t>0:\|V_{t}\|\geq r\} with r>0r>0.

Let the covariance with respect to the average token be defined as Vαxˉ:=m−1∑ν=1mVανV^{\alpha\bar{x}}:=m^{-1}\sum_{\nu=1}^{m}V^{\alpha\nu}, and the average trace be Vˉ:=m−1∑ν=1mVνν\bar{V}:=m^{-1}\sum_{\nu=1}^{m}V^{\nu\nu}. We will need to compute a couple of important moments from the Taylor expansion terms of the Softmax (Lemma C.2)

We are now ready to state our main result.

the diffusion coefficient is defined by Σ(V)=γ2(2−γ2)Σlin(V)+γ4τ0−2[Aαβ,δω]α≤β,δ≤ω\Sigma(V)=\gamma^{2}(2-\gamma^{2})\Sigma_{\text{lin}}(V)+{\gamma^{4}}\tau_{0}^{-2}[\mathcal{A}^{\alpha\beta,\delta\omega}]_{\alpha\leq\beta,\delta\leq\omega}, and

The drift depends on the shaped attention mechanism through S1αδ,βωS_{1}^{\alpha\delta,\beta\omega} and S2αδS_{2}^{\alpha\delta}, the moments of the first and second order terms of the Softmax’s Taylor expansion. On the other hand, the diffusion term depends on the attention solely through S1S_{1}, present in the additional term Aαβ,δω\mathcal{A}^{\alpha\beta,\delta\omega}. The presence of Aαβ,δω\mathcal{A}^{\alpha\beta,\delta\omega} is an intriguing difference compared to shaped ReLU networks, where the diffusion is not affected by the activation function. Both components of the SDE depend on averages over the tokens, reflecting the mixing property of the self-attention mechanism, in which every pair of tokens is compared through dot products to form the attention weights. Finally, notice how the residual branch parameter γ2\gamma^{2} has a dampening effect on the scale of both the drift and the diffusion in a similar way as in fully-connected residual network.

We are now ready to introduce the full shaped Transformer architecture, where we combine the attention and residual layers:

where the coefficients are defined in Theorem 3.2 and Theorem 4.2.

4 On Finite Time Stability of the SDE and Shaped Attention Networks

Although we did not observe numerical instability in majority of our simulations of the shaped attention networks and the corresponding SDE, we did observe that the drift component b(Vt)b(V_{t}) in Theorem 4.2 is cubic in the entries of VtV_{t}. Whenever the drift is not Lipschitz as in this case, we do not have general guarantees for the existence of a solution for all time (see the Feller test for explosions [58, Theorem 5.5.29]). In fact, MLPs with smooth activations also yield non-Lipschitz drift coefficients as seen in .

However, locally Lipschitz coefficients are sufficient to guarantee the existence of local solutions, in the sense of up to a stopping time [59, Proposition 6.9]. Not only does this fact help us establish a precise notion of convergence (Definition 4.1), we can also study the practical implications of this for finite sized attention networks. More specifically, we can inspect the effect of architectural changes to a stopping time.

To demonstrate the potential numerical instabilities, we had to choose an adversarial set of parameters: in particular, an unrealistically large norm (approx. 10n10\sqrt{n}) for the initial tokens X0X_{0}, which enlarges the eigenvalues of V0V_{0} to the order of 100100. Given these initial conditions and a large residual connection weight γ\gamma, we were able to consistently generate numerically unstable behaviour in shaped attention networks (see Figure 5 (left)).

Discussion

Previous work have demonstrated the practical impact scaling limits can have on designing activation functions and tuning hyperparameters . We follow this line of motivations and proposed a novel attention mechanism, which successfully stabilizes the covariance structure in arbitrarily deep Transformers (e.g. Figure 1). The natural next step is to investigate the scaling of gradients in the infinite-depth-and-width limit. As illustrated, the existence of an infinite-width limit for the gradient implies the optimal hyperparameters for the training algorithm will also converge. This type of results allows for tuning of hyperparameters on networks with a much smaller width, yet extends easily to arbitrarily large networks that approximates the same limit, saving massive amounts of computing cost in the process. Given the existence of an infinite-depth-and-width limit for the forward pass, we believe it’s possible to extract optimal hyperparameters from networks with not only a much smaller width, but smaller depth as well.

Preliminary Experiments.

Although this work is primarily theoretical, it is important to consider whether or not the proposed architecture is useful in practice. Given limited computing resources, we chose to only briefly test the feasibility of training the shaped Transformer. Nevertheless, our preliminary experiments show promising results when it comes to training stability. In particular, the shaped Transformer (without LayerNorm) does indeed train at approximately the same speed as well tuned Transformer architectures. Full details of the experiment and results can be found in Appendix D. A more comprehensive set of experiments with different tasks, datasets, and larger networks will be required to confidently determine the practical feasibility of the shaped Transformer, which we defer to future work.

Training Dynamics and Generalization.

As mentioned in the introduction, the limitations of infinite-width NTK theories motivates our study of the proportional infinite-depth-and-width limit. In particular, to address many of the open problems in deep learning theory, we need a faithful and tractable description of training dynamics. Given the results at initialization, the proportional limit holds the potential for such a theory of training as well. Another promising indicator is that deep networks learn features in the proportional regime , which has been identified as a key advantage of neural networks over kernel methods . A precise theory of training will help us understand other types of instabilities during training and improve existing optimization methods. Furthermore, determining the network which training converges to is a necessary step towards a theory of generalization, as demonstrated by the infinite-width approach . In light of our results, we believe that our theory sets the stage for future work on training and generalization in deep learning.

Acknowledgement

CL and ML would like to thank Keiran Paster for insightful discussions. LN would like to thank Sotiris Anagnostidis for support in pre-processing the dataset used for the training experiments of this manuscript. ML is supported by the Ontario Graduate Scholarship and Vector Institute. DMR is supported in part by Canada CIFAR AI Chair funding through the Vector Institute, an NSERC Discovery Grant, Ontario Early Researcher Award, a stipend provided by the Charles Simonyi Endowment, and a New Frontiers in Research Exploration Grant.

References

Appendix A Preliminaries: Covariance SDE Framework

In this section, we will review existing results on Markov chain convergence to an SDE, as well as existing results for the covariance SDE. See also the Appendix of for more details.

In particular, we call the topology that induces this convergence the Skorohod topology [68, Theorem A5.3].

On a heuristic level, if we have sequences of Markov chains YnY^{n} that satisfy the following type of Euler updates

Now we will state the main technical result in this section.

Suppose otherwise bn,σnb_{n},\sigma_{n} are only locally Lipschitz (but still uniform in nn), then XnX^{n} converges locally to XX in the same topology (see Definition 4.1). More precisely, for any fixed r>0r>0, we consider the stopping times

then the stopped process Xt∧τnnX^{n}_{t\wedge\tau^{n}} converges in distribution to the stopped solution Xt∧τX_{t\wedge\tau} of the above SDE in the same topology.

We will briefly recall the main result of next. As mentioned earlier in the text, the setting is for MLPs defined as follows

where the coefficients were given in Theorem 3.2.

Appendix B SDE for Residual Network

Recall that we adopt the following model:

It is easy to see that the cross product terms have (conditional) mean zero. For the term with γ2\gamma^{2}:

showed that if the non linearity scaling exponent is p=1/2p=1/2, then:

where ν(ρ)=(c++c−)22π(1−ρ2+ρarccos(ρ))\nu(\rho)=\frac{(c_{+}+c_{-})^{2}}{2\pi}\left(\sqrt{1-\rho^{2}}+\rho\text{arccos}(\rho)\right). Using this result, and summing and subtracting the mean of T2\mathcal{T}_{2}:

Furthermore, we need second order moments of T1\mathcal{T}_{1} and T2\mathcal{T}_{2}. We derive them in the following Lemmas:

Recall the definition of T1αβ\mathcal{T}_{1}^{\alpha\beta}:

Recall the definition of T2αβ\mathcal{T}_{2}^{\alpha\beta}:

For the mean, using the independence between WpreW^{\text{pre}} and WpostW^{\text{post}}, we have that:

which is the desired result. The final expression is the result of the aforementioned expansion for K1K_{1} . For the covariance, we have that:

Let’s look at each term of the sum separately. For the first term we have that:

The other two summands can be solved similarly:

where in the final step we have used the fact that c=1+O(n−1/2)c=1+\mathcal{O}(n^{-1/2}). This completes the proof. ∎

Finally we also have that T1\mathcal{T}_{1} and T2\mathcal{T}_{2} are uncorrelated, i.e.

Using λ2+γ2=1\lambda^{2}+\gamma^{2}=1, and summing and subtracting the mean of T2\mathcal{T}_{2}, we have that:

Using Lemma B.2 for the mean of T2\mathcal{T}_{2}, we have that:

Finally, applying Proposition A.2, we get the desired result. ∎

Appendix C SDE for Softmax-based Attention

Dot-product attention applies the Softmax row-wise to the following matrix:

For the second moment we have the following Lemma.

C.2 Shaping the Softmax

where τ−1\tau^{-1} is a temperature parameter that regulates the entropy of the resulting categorical distribution; τ\tau is often chosen to be large (scale as a power of nn or dd), as low temperature results in unstable training.

To transform the Softmax into the desired form of Aαβ=δαβ+O(n−1)A^{\alpha\beta}=\delta_{\alpha\beta}+\mathcal{O}(n^{-1}), we first center the attention matrix around the identity matrix:

and then examine the Taylor-expansion of the Softmax w.r.t τ−1\tau^{-1}:

which admits the following Taylor-expansion:

where Yα‾:=1m∑νYαν\overline{Y^{\alpha}}:=\frac{1}{m}\sum_{\nu}Y^{\alpha\nu}, and (Yα)2‾:=1m∑ν(Yαν)2\overline{(Y^{\alpha})^{2}}:=\frac{1}{m}\sum_{\nu}(Y^{\alpha\nu})^{2}.

C.3 Lemmas on Moments of Shaped Attention

where Vˉ=1m∑νVνν\bar{V}=\frac{1}{m}\sum_{\nu}V^{\nu\nu} and xˉ=1m∑νxν\bar{x}=\frac{1}{m}\sum_{\nu}x^{\nu} is the average token.

Using Lemma C.1 and linearity of expectation:

where Vˉ=1m∑βVββ\bar{V}=\frac{1}{m}\sum_{\beta}V^{\beta\beta}. ∎

Note that the terms above can be computed using Lemma C.2.

The results is an immediate consequence of Lemma C.4 and Lemma C.3, where only the terms that do not cancel out are kept. ∎

C.4 Neural Covariance SDE for Stable Attention

We need to compute the moments for these two quantities:

Hence, in expectation with respect to WW:

An identical argument can be made for the remaining three summands. Hence, taking expectation with respect to the Softmax weights:

Taking expectation w.r.t the Softmax weights, we get the desired result.

For second moment, we can take the conditional expectation:

By taking expectation w.r.t the Softmax parameters, we get the desired result. ∎

Now we can use Lemma C.3 and Lemma C.4 to compute the moments of AA. For the second summand, we simply have:

For the first summand, recall from Lemma C.5 that:

We are now ready to re-state and proof of Theorem 4.2.

Plugging in the expression, and using λ2+γ2=1\lambda^{2}+\gamma^{2}=1, we have that the drift is:

From here, it is evident that in order to have the drift scaling as O(1/n)\mathcal{O}(1/n) we need to choose:

Covariance.

Furthermore, we have set λ2+γ2=1\lambda^{2}+\gamma^{2}=1 and τ2=τ02nnk\tau^{2}=\tau_{0}^{2}nn_{k} to have the right scaling for the drift.

Using Lemma C.6 and Lemma C.8, we have that:

Now we can apply Proposition A.2 for locally Lipschitz drift and covariance coefficients, which gives us the desired result in local convergence in the Skorohod topology.

We will also restate and prove Corollary 4.3.

To combine the results of Theorem 3.2 and Theorem 4.2, it is sufficient to combine the following (simplified) iterated Markov updates into one Markov chain

Since in the limit, we have that either updates are infinitesimal, i.e.

which converges to the following SDE with two Brownian motions using Proposition A.2

Observe that since the two Brownian motions Bt,Bt′B_{t},B_{t}^{\prime} are independent, it’s equivalent to write

We recover the desired results from considering a more general form of the iterated Markov updates as in Proposition A.2, which do not hinder the above derivation.

Appendix D Preliminary Experiments

and propose two ways to set γ1,γ2\gamma_{1},\gamma_{2} during training. In both alternatives we initialize γ1,γ2=1\gamma_{1},\gamma_{2}=1, thus leveraging the stability properties of shaped attention at initialization. During training, we either:

Recover. Linearly decrease γ1,γ2\gamma_{1},\gamma_{2} to zero with in the first 40004000 steps, thus recovering the standard attention layer. Apply the same schedule for the shaped-ReLU slope s−s_{-}, recovering the usual ReLU activation. This approach recovers the vanilla Transformer architecture (without LayerNorm).

Learn. Learn all the shaping parameters γ1\gamma_{1}, γ2\gamma_{2} and s−s_{-}.

The intuitive rationale behind these choices is that at initialization we want good signal propagation and a non-degenerate covariance (according to our theory, this requires the shaped attention and ReLU). On the other hand, we also allow the model to more dynamically make use of the nonlinearity during training to modify the correlation structure with Recover or Learn. We report that without either adjustment, shaped attention is still trainable but at much slower rates.

We also incorporate the 1n\frac{1}{\sqrt{n}} factor into the initialization of the queries and keys weights by decreasing the variance of their entries by a factor of 1n\frac{1}{n}. This allows us to not re-tune the learning rate for the queries and keys. We stress that at at initialization, the two formulations (τ=nnk\tau=\sqrt{nn_{k}} and τ=nk\tau=\sqrt{n_{k}} with decreased weight’s variance) are equivalent. In both alternatives we train all the skip connection parameters λ\lambda, γ\gamma, and initialize γ\gamma in the grid (0.05,0.1,0.2)(0.05,0.1,0.2) and set τ0=1\tau_{0}=1. We report training instabilities (loss divergence) for larger values of γ\gamma. All models —including the baselines — use Adam with learning rate warmup of 40004000 steps, the learning rate is tuned in the grid (0.0001,0.0005,0.001,0.005)(0.0001,0.0005,0.001,0.005). We report the train/test loss after 100K100K optimization steps, averaged over 44 random seeds. All the other experimental details can be found in Section D.1.

In Table 1, we compare the train/test loss of the two variant of shaped attention with the baseline Pre-LN model. Notice that our model (in both variants) achieves comparable performance to standard Transformers across all the reported values of γ\gamma.

GLUE evaluation

Furthermore, we evaluate the trained models on three datasets from the GLUE benchmark and summarize the results in Table 2. These demonstrate that our shaped Transformers holds promise by outperforming our pre-ln baseline.

Entropy Collapse for Large Learning rates.

To understand the sources of training instability, we keep track of the entropy of the probability distribution induced by the Softmax, as it has been observed that the Transformer’s training is unstable in the low-entropy regime . The entropy is calculated for each row of the Softmax matrix, and it is averaged across rows and heads. The results are in Fig. 6. Notice how for the large learning rate regime observed in Fig. 6, the entropy collapses for the baseline model, but not for the proposed shaped Transformer. Entropy collapse indicates that the Softmax distribution degenerates to a point mass, which is itself caused by large logits. Remarkably, this phenomenon does not affect the recover setting, despite recovering the Transformer architecture (without layer normalization) after the warm-up period.

D.1 Experimental details

We use a subset of the English Wikipedia 20220301.en and English bookcorpus datasets . The sentences are tokenized using by pre-training a tokenizer on the training set. We use a vocabulary size of 3200032000, and a maximum sequence length of 128128 tokens.

Model parameters.

We use an embedding size of n=768n=768 and 88 multi-attention heads. The batch size is fixed to 3232 sequences. All the initial weights are sampled from N(0,n−1)\mathcal{N}(0,n^{-1}), with the exception of the queries and keys’ weights WK,WQW^{K},W^{Q} in the shaped attention case, that are sampled from N(0,n−3/2)\mathcal{N}(0,n^{-3/2}) (as explained in Appendix D). The feedforward layer maps the n=768n=768-dimensional embedding to the larger dimension 30723072, as in the Hugging face implementation of Bert .

Optimization.

We train using Adam with betas parameters (0.90.9, 0.9990.999) and learning rate chosen in the grid (0.0001,0.0005,0.001,0.005)(0.0001,0.0005,0.001,0.005). We do not use weight decay.

Computational Resources.

The experiments are executed on Nvidia DGX-1 GPU nodes equipped with 4 20-core Xeon E5-2698v4 processors, 512 GB of memory and 8 Nvidia V100 GPUs.

Appendix E Additional Figures