Infinite attention: NNGP and NTK for deep attention networks

Jiri Hron, Yasaman Bahri, Jascha Sohl-Dickstein, Roman Novak

Introduction

One of the currently most active research directions in theoretical deep learning is the study of NN behaviour as the number of parameters in each layer goes to infinity (e.g., Matthews et al., 2018; Lee et al., 2018; Garriga-Alonso et al., 2019; Novak et al., 2019; Li & Liang, 2018; Allen-Zhu et al., 2019; Du et al., 2019; Arora et al., 2019; Yang, 2019b). Building upon these efforts, we study the asymptotic behaviour of NNs with attention layers (Bahdanau et al., 2015; Vaswani et al., 2017) and derive the corresponding neural network Gaussian proccess (NNGP) and Neural Tangent kernels (NTK, Jacot et al., 2018; Lee et al., 2019).

Beyond their recent empirical successes (e.g., Radford et al., 2019; Devlin et al., 2019), attention layers are also interesting from the theoretical perspective as the standard proof techniques used to establish asymptotic Gaussianity of the input-to-output mappings represented by wide NNs (Matthews et al., 2018; Yang, 2019b) cannot be applied.

where ζ\zeta is the row-wise softmax function.

Now observe that dim⁡G(x)=ds×ds\dim G(x)=d^{s}\times d^{s} where the spatial dimension dsd^{s} stays finite even as the number of parameters—here proportional to dd—goes to infinity. As we will show rigorously in Section 3, this fact combined with the d−1/2d^{-1/2} scaling causes each column of f(x)f(x) to be a linear combination of the same stochastic matrix ζ(G(x))\zeta(G(x)), and thus statistically dependent even in the infinite width limit.

Since the exchangeability based arguments (Matthews et al., 2018; Garriga-Alonso et al., 2019) require that certain moment statistics of f(x)f(x) asymptotically behave as if its columns were independent (see condition b in lemma 10, Matthews et al., 2018), they do not extend to attention layers in a straightforward manner. Similarly, the proofs based on Gaussian conditioning (Novak et al., 2019; Yang, 2019b) require that given the input xx, the conditional covariance of each column of f(x)f(x) converges (in probability) to the same deterministic positive semidefinite matrix (see propositions 5.5 and G.4 in Yang, 2019b) which will not be the case due to the aforementioned stochasticity of ζ(G(x))\zeta(G(x)).

Among the many interesting contributions in (Yang, 2019b), the author proposes to resolve the above issue by replacing the d−1/2d^{-1/2} scaling in Equation 1 by d−1d^{-1} which does enable application of the Gaussian conditioning type arguments. However, it also forces the attention layer to only perform computation similar to average pooling in the infinite width limit, and reduces the overall expressivity of attention even if suitable modifications preventing the pooling behaviour are considered (see Section 3.2).

We address the above issues by modifying the exchangability based technique and provide a rigorous characterisation of the infinite width behaviour under both the d−1/2d^{-1/2} and d−1d^{-1} scalings. We also show that positional encodings (Gehring et al., 2017; Vaswani et al., 2017) can improve empirical performance even in the infinite width limit, and propose modifications to the attention mechanism which results in further gains for both finite and infinite NNs. In experiments, we moderately improve upon the previous state-of-the-art result on CIFAR-10 for GP models without data augmentation and advanced preprocessing (cf. Yu et al., 2020). Finally, since attention is often applied to text datasets, we release code allowing applications of NNGP/NTK models to variable-length sequences, including an example on the IMDb reviews dataset.

Definitions and notation

Attention and Gaussian process behaviour

Throughout the rest of this paper, we restrict our focus to increasingly wide NNs including at least one attention layer. In particular, we consider sequences of NNs such that

and the reader should thus interpret any statements involving n→∞n\to\infty as implicitly assuming Equation 4 holds.

As illustrated in Figure 1, use of the d−1/2d^{-1/2} scaling within a single-head architecture leads to a scale mixture behaviour of the attention layer outputs as the number of parameters goes to infinity. To obtain a Gaussian limit, Yang (2019b, appendix A) proposes to replace the definition in Equation 2 by Gn(x)=(dnG)−1Qn(x)Kn(x)⊤G_{n}(x)=(d_{n}^{G})^{-1}Q_{n}(x)K_{n}(x)^{\top}, i.e., the use of d−1d^{-1} scaling. The desired result then follows:

Under the d−1d^{-1} scaling and the assumptions stated in (Yang, 2019b):

fnf_{n} converges in distribution to f∼GP(0,κ)f\sim\mathcal{GP}(0,\kappa) with

and f⋅kf_{\cdot k} and f⋅lf_{\cdot l} are independent for any k≠lk\neq l.

An analogous result also holds for multi-head attention architectures which follows by the usual argument for fully connected layers as long as either the number of embedding dimensions per head or the number of heads goes to infinity.

While Theorem 1 is a good starting point, several issues have to be resolved before using the kernel function described in Equation 6 in practice. Firstly, since WnQW_{n}^{Q} and WnKW_{n}^{K} are initialised independently, the d−1d^{-1} scaled inner products of keys and queries will converge to zero (the mean), and thus for any a,ia,i and xx, ζˉaix→(ds)−1\bar{\zeta}_{ai}^{x}\to(d^{s})^{-1} in probability by the continuous mapping theorem. This issue was already noted by Yang in appendix A but not discussed further as the main focus of the paper lies elsewhere. In any case, substituting (ds)−1(d^{s})^{-1} for all the ζˉ\bar{\zeta} coefficients will make κab (x,x′)=κij (x,x′)\kappa_{ab}\,(x,x^{\prime})=\kappa_{ij}\,(x,x^{\prime}) for any a,b,i,j∈[ds]a,b,i,j\in[d^{s}], and in fact all of these entries will be equivalent to output of a simple global average pooling kernel (Novak et al., 2019, equation 17).In fact, the asymptotic distribution induced by such an attention layer followed by flatten and dense layers is the same as that induced by global average pooling followed by a dense layer.

Perhaps the simplest way to address the above issue is by drawing the initial weights such that WnQ=WnKW_{n}^{Q}=W_{n}^{K}. This will ensure that the key and query for a particular spatial dimension will point in the same direction and thus the attention weight corresponding to itself will be large with high probability. The resulting formula for κab (x,x′)\kappa_{ab}\,(x,x^{\prime}) is

Since Equation 7 resolves the issue of reduction to average pooling, a natural question is whether swapping d−1/2d^{-1/2} for d−1d^{-1} has any undesirable consequences in the infinite width limit. As we will see, this question can be answered in affirmative. In particular, we start by a proposition inspired by (Cordonnier et al., 2020) in which the authors show that an attention layer with a sufficient number of heads is at least as expressive as a standard convolutional layer, and that attention layers often empirically learn to perform computation akin to convolution. In contrast, Proposition 2 proves that there is no initial distribution of WnQW_{n}^{Q} and WnKW_{n}^{K} which would recover the convolutional kernel (Novak et al., 2019; Garriga-Alonso et al., 2019) in the infinite width limit.

where dfd_{f} is the dimension of the (flattened) convolutional filter, Na,Nb⊂[ds]N_{a},N_{b}\subset[d^{s}] are the ordered subsets of pixels which are used to compute the new values of pixels aa and bb, respectively, and Na(i),Nb(i)N_{a}(i),N_{b}(i) are the iith pixels in Na,NbN_{a},N_{b}.

In the next section, we will see that the convolutional kernel can be recovered under the d−1/2d^{-1/2} scaling (Proposition 4). However, we need to establish convergence scaling first.

As discussed in Section 1, single-head attention architectures can exhibit non-Gaussian asymptotic behaviour under the d−1/2d^{-1/2} scaling. This is inconvenient for our purposes as many modern NN architectures combine attention with fully connected, convolutional, and other layer types, all of which have Gaussian NNGP and NTK limits (e.g., Novak et al., 2019; Garriga-Alonso et al., 2019; Yang, 2019b). This Gaussianity simplifies derivation of the infinite width behaviour of many architectures and allows for easy integration with existing software libraries (Novak et al., 2020). Fortunately, the output of an attention layer becomes asymptotically Gaussian when the number of heads becomes large.

We can now revisit our argument from the previous section, and prove that unlike in Proposition 2, d−1/2d^{-1/2} scaling ensures a convolutional kernel can in principle be recovered.

Under the d−1/2d^{-1/2} scaling, there exists a distribution over GG such that for any x,x′x,x^{\prime} and a,b,i,ja,b,i,j

Beyond the vanilla attention definition

Figure 3 shows the results across varying hyperparameters and random seeds, and Table 8 (Section A.1.3) reports accuracies attained under optimal hyperparameter settings. As you can see, both the replacement of softmax and addition of layer normalisation significantly increases the performance of the NN, with ζ(x)=x\zeta(x)=x and at_output normalisation being the best across variety of hyperparameter choices.

In light of the above, we will restrict our attention to the identity function alternative for ζ\zeta in the rest of the paper, and contrast its performance with the standard softmax choice where possible (finite NNs, and infinite attention NNs under the d−1d^{-1} scaling—see Theorem 1). Similarly, we will also leverage the at_output layer normalisation over the embedding dimension in our experiments. As shown by Yang (2019b, appendix A), layer normalisation does not prevent Gaussianity of the infinite width limit (see Table 1 for the associated NNGP and NTK kernel transformations).

2 Positional encodings

While substituting the identity function for ζ\zeta as suggested in Section 4.1 would technically allow us to move on to the experimental evaluation already, we found that positional encodings are as important in the infinite width limit as they are for the finite attention layers (Vaswani et al., 2017). Since there are many possible variants of the positional encoding implementation, we focus only on the major points here and provide more detail in Appendix C.

2.2 Structured positional encodings

As mentioned, the main purpose of positional encodings is to inject structural information present in the inputs which would be otherwise ignored by the attention layer. A natural way to resolve the issues discussed in previous section is thus to try to incorporate similar information directly into the RR covariance matrix. In particular, we propose

where ρ,φ>0\rho,\varphi>0 are hyperparameters, rh(a,b)r_{h}(a,b) and rv(a,b)r_{v}(a,b) are the absolute horizontal and vertical distances between the pixels aa and bb divided by the image width and height respectively, and rs(a,b)r_{s}(a,b) is the absolute distance between the relative position of tokens aa and bb, e.g., if aa is the 4th token out of 7 in the first, and bb is the 2nd token out of 9 in the second string, then rs(a,b)=∣47−29∣r_{s}(a,b)=|\frac{4}{7}-\frac{2}{9}|.

The above reasoning only provides the motivation for modifying the attention weights using positional encodings but not necessarily for modifying the asymptotic distribution of the values VV. Adding positional encodings only inside the ζ\zeta is not uncommon (e.g., Shaw et al., 2018), and thus we will also experiment with kernels induced by adding positional encodings only to the inputs of QnQ_{n} and KnK_{n}, leading to

under the d−1d^{-1} scaling (cf. Equation 12), and

under the d−1/2d^{-1/2} scaling (cf. Equation 13).

Experiments

We evaluate the attention NNGP/NTK kernels on the CIFAR-10 (Krizhevsky, 2009) and IMDb reviews (Maas et al., 2011) datasets. While IMDb is a more typical setting for attention models (Section 5.2), we included CIFAR-10 experiments (Section 5.1) due to desire to compare with other NNGPs/NTKs on an established benchmark (e.g., Novak et al., 2019; Du et al., 2019; Yu et al., 2020), and the recent successes of attention on vision tasks (e.g., Wang et al., 2017, 2018; Hu et al., 2018; Woo et al., 2018; Chen et al., 2018; Ramachandran et al., 2019; Bello et al., 2019). Our experimental code utilises the JAX (Bradbury et al., 2018) and Neural Tangents (Novak et al., 2020) libraries.

We have run two types of experiments on CIFAR-10: (i) smaller scale experiments focused on understanding how different hyperparameters of the attention kernel affect empirical performance; (ii) a larger scale experiment comparing attention kernels to existing NNGP/NTK benchmarks. The smaller scale experiments were run on a randomly selected subset of six thousand observations from the training set, with the 2K/4K train/validation split. This subset was used in Figures 2 and 4, and for hyperparameter tuning. Selected hyperparameters were then employed in the larger scale experiment with the usual 50K/10K train/test split.

All kernels evaluated in this section correspond to NN architectures composed of multiple stacked convolutional layers with ReLU activations, followed by either simple flattening, global average pooling (GAP), or one of our attention kernels itself followed by flattening and, except for the Vanilla attention case (see Table 1), also by layer normalisation; the output is then computed by a single dense layer placed on top. The choice to use only one attention layer was made to facilitate comparison with (Novak et al., 2019; Du et al., 2019; Yu et al., 2020) where the same set-up with a stack of convolutional layers was considered. Adding more attention layers did not result in significant gains during hyperparameter search though. Exact details regarding data normalisation, hyperparameter tuning, and other experimental settings can be found in Appendix A.

The most important observations from the smaller scale experiments are captured in Figure 4 which shows the validation accuracy of various NNGP models as a function of kernel choice and number of convolutional layers (depth) preceding the final flatten/GAP/attention plus dense block. Firstly, notice that except for the Flatten model, all other kernel choices achieve their best performance at smaller depths which is consistent with existing literature (Arora et al., 2019; Yu et al., 2020).

Secondly, observe that both the Struct and Residual attention kernels significantly outperform the Vanilla one, demonstrating that the use of positional embeddings and layer normalisation is helpful even in the infinite width limit as claimed in Section 4.2. In contrast, we did not find significant evidence for ζ(x)=x\zeta(x)=x outperforming the standard softmax choice as was the case for finite networks (see Figure 3), with the best set of hyperparameters for Struct d−1d^{-1} with softmax being only marginally better than the best results with the identity function (recall that no d−1/2d^{-1/2} kernels use ζ=softmax\zeta=\text{softmax} due to the intractability discussed in Section 4). This finding provides hope that the d−1/2d^{-1/2} kernels also do not sacrifice much in terms of performance by using identity for ζ\zeta, but also points to salient differences between the qualitative effects of individual hyperparameter choices in finite and infinite attention layers.

Using the insights from the smaller scale experiments, we ran the larger scale experiment on the full dataset using eight layer models and the Struct and Residual attention kernels. We used the positional embedding covariance matrix defined in Equation 14 in both cases, and d−1d^{-1} with softmax for the Struct kernel (further details in Section A.1.5). The results can be found in Table 2. As you can see, attention performs significantly better than the GAP kernel (Arora et al., 2019), and also provides a moderate improvement over the recent local average pooling (LAP) results (Yu et al., 2020). Since we used the validation accuracy from smaller scale experiments to determine our hyperparameters, we are comparing against the best cross-validation results from (Yu et al., 2020) for fairness.

2 IMDb reviews

Although there has been interest in applying attention in vision, to date it has been predominantly recognized for performance on language tasks. However, most of available NNGP/NTK kernel implementations (Matthews et al., 2018; Lee et al., 2018; Garriga-Alonso et al., 2019; Arora et al., 2019; Yang, 2019b; Yu et al., 2020) are hard-coded for the specific experiments performed in the respective paper. Neural Tangents (Novak et al., 2020) allows for some flexibility, yet still accepts only inputs of fixed length and having exactly zero (i.e. inputs to fully connected networks) or two (images for CNNs) spatial dimensions.

We release code allowing use of NNGP/NTK models (with or without attention) on inputs of variable spatial extent and arbitrary dimensionality (e.g., one spatial dimension for texts and time series, three spatial dimensions for videos). Our implementation seamlessly extends the Neural Tangents library, enabling research and application of NNGP and NTK models to new domains with almost no extra effort.

As an example, we present the first benchmarks of simple NNGP and NTK models on the IMDb sentiment classification dataset in Table 3. We observe that Struct kernels outperform the GAP-only kernel (corresponding to linear regression on the word embeddings mean), but provides marginal benefit compared to a fully connected model on top of the pooling layer (GAP-FCN). We conjecture this is due to high-quality word embeddings partially incorporating the inductive bias of the considered model. Indeed, we further demonstrate this effect by contrasting the gaps in performance between different kernel families on high- and low-quality word embeddings in Table 4.

Naturally, our sample IMDb results are not competitive with the state-of-the-art, which achieve up to 97.4% (Thongtan & Phienthrakul, 2019, Table 4). However, we hope they will be a useful baseline for future research in infinite width sequence models, and that our codebase will substantially facilitate the process by enabling variable-length, arbitrary-dimensional input processing.

Conclusion

Unlike under the d−1d^{-1} scaling of Q(x)K(x)⊤Q(x)K(x)^{\top} proposed in (Yang, 2019b), the standard d−1/2d^{-1/2} scaling may lead to non-Gaussian asymptotic behaviour of attention layer outputs. Gaussianity of the limit can however be obtained by taking the number of heads to infinity. We explored the effect of positional encodings and replacements for the softmax function in attention layers, leading to improved performance for both finite and infinite attention architectures. On CIFAR-10, attention improves moderately upon the previous state-of-the-art for GPs without trainable kernels and advanced data preprocessing (Yu et al., 2020). We further released code allowing application of NNGP/NTK kernels to variable-length sequences and demonstrated its use on the IMDb reviews dataset. While caution is needed in extrapolation of any results, we hope that particularly Figure 3 and Table 2 inspire novel NN architectures and kernel designs.

Acknowledgements

We thank Jaehoon Lee for frequent discussion, help with scaling up the experiments, and feedback on the manuscript. We thank Prajit Ramachandran for frequent discussion about attention architectures. We thank Greg Yang, Niki Parmar, and Ashish Vaswani, for useful discussion and feedback on the project. Finally, we thank Sam Schoenholz for insightful code reviews.

References

Appendix A Experimental details

The CIFAR-10 datasest (Krizhevsky, 2009) was fetched using the TensorFlow datasetshttps://www.tensorflow.org/datasets/catalog/cifar10.

In all of the CIFAR-10 experiments, the data was preprocessed by subtracting mean and dividing by a standard deviation for each pixel and data point separately (equivalent to using LayerNorm as the first layer). We inflated all of the standard deviations by 10−1510^{-15} to avoid division by zero.

All the classification tasks were converted into regression tasks by encoding the targets as CC–dimensional vectors, where CC is the number of classes, with the entry corresponding to the correct label set to C−1C\frac{C-1}{C} and all other entries to −1C-\frac{1}{C}. This enabled us to perform closed form NNGP and NTK inference using the Gaussian likelihood/MSE loss.

The hyperparameter search was on a fixed architecture with 8x Convolution + ReLU, Attention, Flatten, and a Dense readout layer. We used 1.75621.7562 and 0.18410.1841 respectively for the weight and bias variances as in (Novak et al., 2019, appendix G.1) except for the attention output variance σO2\sigma_{O}^{2} which was set to one. The convolutional layers were used with the SAME padding, stride one, and filter size 3×33\times 3. For attention kernels with positional encodings, the reported ρ\rho parameter (Equation 14) is actually ρ/(σQ2σK2)\rho/(\sigma_{Q}^{2}\sigma_{K}^{2}) so that the relative scale of the contribution of RR remains the same with changing σQ2σK2\sigma_{Q}^{2}\sigma_{K}^{2}.

There were two stages of the hyperparameter search, first to identify the most promising candidates (Table 5), and second to refine the parameters of these candidate kernels (Table 6). The second stage also included the residual attention kernel (Equation 16); the α\alpha in the second table should thus be interpreted as the one stated in Equation 16 (cf. Appendix D). The best hyperparameters used in Figure 4 and Table 2 can be found in a bold typeset in Table 6.

All computation was done in 32-bit precision, and run on up to 8 NVIDIA V100 GPUs with 16Gb of RAM each.

A.1.2 Details for Figure 2

The downsampling was performed using skimage.transform.resize with parameters mode="reflect" and anti_aliasing=True, using downsampled height and width of size 8 as mentioned.

Both the convergence and accuracy plots are for the d−1/2d^{-1/2} vanilla NNGP kernel with ζ=softmax\zeta=\text{softmax}. The intractable softmax integral of the limiting covariance function was estimated using MC integration with 2048 samples.

We used 1.75621.7562 and 0.18410.1841 respectively for the weight and bias variances as in (Novak et al., 2019, appendix G.1) for all the convolutional and dense layers, 1.75621.7562 for the σK2,σQ2\sigma_{K}^{2},\sigma_{Q}^{2} and σV2\sigma_{V}^{2}, and σO2=1\sigma_{O}^{2}=1. The convolutional layer used VALID paddingstride one, and filter size 3×33\times 3.

As in (Novak et al., 2019), The reported distance between kernel matrices is the logarithm of

where K^\hat{\mathcal{K}} and K\mathcal{K} are respectively the empirical and the predicted theoretical covariance matrices for the training set.

All computation was done in 32-bit precision, and run on up to 8 NVIDIA V100 GPUs with 16Gb of RAM each.

A.1.3 Details for Figure 3

We used a 45K/5K train/validation split of the usual 50K CIFAR-10 training set and reported the validation set accuracy after training for 1000 epochs with batch size 64 and the Adam optimiser.

The attention layers used the usual d−1/2d^{-1/2} scaling of the query/key inner products, and the convolutional layers used the SAME padding, stride one, and filter size 3×33\times 3. We used 2.02.0 and 10−210^{-2} respectively for the weight and bias variances except in the attention where σQ2=σK2=σV2=2\sigma_{Q}^{2}=\sigma_{K}^{2}=\sigma_{V}^{2}=2 but σO2=1\sigma_{O}^{2}=1. Further, we used the append type positional encodings (Section 4.2) with the same embedding dimension as n_channels (Table 7), thus doubling the embedding dimension of the attention layer inputs.

All computation was done in 32-bit precision, and run on a single NVIDIA V100 GPU with 16Gb of RAM each.

A.1.4 Details for Figure 4

We used 1.75621.7562 and 0.18410.1841 respectively for the weight and bias variances as in (Novak et al., 2019, appendix G.1) except for the attention output variance σO2\sigma_{O}^{2} which was set to one. The convolutional layers were used with the SAME padding, stride one, and filter size 3×33\times 3. For the vanilla attention kernels, we report the best performance over σQσK={10−3,10−1,1,2,10}\sigma_{Q}\sigma_{K}=\{10^{-3},10^{-1},1,2,10\} at each depth. The Struct and Residual were used with the best hyperparameters found during hyperparameter search as reported in Section A.1.1.

All computation was done in 32-bit precision, and run on up to 8 NVIDIA V100 GPUs with 16Gb of RAM each.

A.1.5 Details for Table 2

The best set-up from Section A.1.1 was used (including the best hyperparameters as stated in Table 6).

All computation was done in 64-bit precision, and run on up to 8 NVIDIA V100 GPUs with 16Gb of RAM each.

A.2 IMDb

The IMDb reviews dataset (Maas et al., 2011) was fetched using TensorFlow datasetshttps://www.tensorflow.org/datasets/catalog/imdb_reviews.

All sentences were truncated or padded to 1000 tokens using the default settings of tf.keras.preprocessing.text.Tokenizerhttps://www.tensorflow.org/api_docs/python/tf/keras/preprocessing/text/Tokenizer. No words were removed from the embedding model dictionary. Tokens were embedded using GloVe embeddings (Pennington et al., 2014) with no other pre-processing. Binary targets were mapped to {−0.5,0.5}\left\{-0.5,0.5\right\} values. Diagonal regularizers for inference were selected based on validation performance among the values of 10−7,10−6,…,110^{-7},10^{-6},\dots,1 multiplied by the mean trace of the kernel.

When applicable, all models used ReLU nonlinearities, Struct (Structured positional encoding, d−1d^{-1} scaling, Table 1) kernel with ζ\zeta being the row-wise softmax function (Equation 20), decaying positional embeddings used only for the attention keys and queries, with φ=2.5\varphi=2.5 (Equation 14), α=0.75\alpha=0.75, and ρ=1\rho=1 (Equation 11). These parameters were selected based on preliminary experiments with CIFAR-10, and fine-tuning on IMDb specifically is an interesting avenue for future research.

All preliminary and validation experiments were carried out in 32-bit precision, while test evaluation (reported in the Table 3 and Table 4) were done in 64-bit precision. All experiments were run on machines with up to 8 NVIDIA V100 GPUs with 16Gb of RAM each.

A.2.2 Details for Table 3

Words were embedded using GloVe 840B.300d embeddings.

The embedding model was selected on a small-scale experiment (4000 train and 4000 validation sets) among GloVe 6B 50-, 100-, 200-, and 300-dimensional variants, as well as GloVe 840B.300d, and 1024-dimensional ELMO (Peters et al., 2018) embeddings (using TensorFlow Hubhttps://tfhub.dev/google/elmo/3). In this preliminary experiment, GloVe 840B.300d, GloVe6B.300d, and ELMO.1024d performed similarly, and GloVe 840B.300d was chosen for the full dataset experiment.

The validation experiment was run on the 25K training set partitioned into a 15K and 10K training and validation sets, with the best models then evaluated on the 25K training and 25K test sets.Precisely, subsets of sizes 14880/9920 and 24960/24960 were used to make the dataset be divisible by 8 (the number of GPUs) times 20 (the batch size), which is a technical limitation of the Neural Tangents (Novak et al., 2020) library.

All layers used weight and bias variances 22 and 0.010.01 respectively, expect for attention outputs and values variances which were set to 11, and the top linear readout layer with weight variance 1 and no bias.

GAP-only, doing only global average pooling over inputs followed by the linear readout.

GAP-FCN, in which GAP was followed by 0, 1, or 2 fully connected layers.

Struct, allowing the same models as GAP-FCN, except for necessarily having an attention layer before GAP.

Each class could also have an optional LayerNorm layer following GAP. The best model from each class was then evaluated on the test set.

A.2.3 Details for Table 4

All convolutional layers used the total window (context) size of 9 tokens, stride 1, and SAME (zero) padding.

Experiments were run on a 3200/1600/1600 train/validation/test splits. Four classes of models were considered:

GAP-only, identical to the one in Section A.2.2.

GAP-FCN, also identical to the one in Section A.2.2.

CNN-GAP, allowing the same models as in GAP-FCN, but having GAP preceeded by 0, 1, 2, 4, or 8 CNN layers.

Struct, allowing the same models as in CNN-GAP, but having 1 or 2 attention layers (each optionally followed by LayerNorm over channels) before GAP. If the model also had CNN layers, attention and CNN layers were interleaved, attention layers being located closer to GAP (for example, a model with 8 CNN layers and 2 attention layers would have 7 CNN layers followed by attention, CNN, attention, GAP).

All models were allowed to have either ReLU or Erf nonlinearity, with weight and bias variances set to 2 and 0.01 for ReLU, and 1.7562 and 0.1841 for Erf, with the same values used by attention keys and queries layers, but having variance 1 for values and output layers. The readout linear layer had weight variance 1 and no bias.

Appendix B Proofs

as defined on the same probability space, and thus allows us to make claims about convergence in probability and similar.

Finally, we will be using the NTK parametrisation (Jacot et al., 2018) within the NTK convergence proofs, i.e., we implicitly treat each weight Wij∼N(0,σ2/d)W_{ij}\sim\mathcal{N}(0,\sigma^{2}/d), i.i.d., as W=σdW~W=\frac{\sigma}{\sqrt{d}}\widetilde{W} where only W~\widetilde{W} is trainable. This parametrisation ensures that not only the forward but also the backward pass are properly normalised; under certain conditions, proofs for NTK parametrisation can be extended to standard parametrisation (Lee et al., 2019).

We are now prepared to apply lemma 10 from (Matthews et al., 2018) which we restate (with minor modifications) here.

Then Sn⇝ZS_{n}\rightsquigarrow Z, where Z=0Z=0 (a.s.) if σ∗2=0\sigma^{2}_{*}=0, and Z∼N(0,σ∗2)Z\sim\mathcal{N}(0,\sigma^{2}_{*}) otherwise.

Substituting Sn=TnS_{n}=\mathcal{T}_{n} and Xn,h=γn,hX_{n,h}=\gamma_{n,h}, convergence of Tn\mathcal{T}_{n} follows from Lemma 7:

Exchangeability requirement is satisfied by Lemma 8.

Zero mean and covariance follow from Lemma 9.

Convergence of variance is established in Lemma 10.

Combining the above with Lemmas 33 and 12 concludes the proof. ∎

γn,h\gamma_{n,h} are exchangeable over the index hh.

suggests the desired result could be obtained by application of Theorem 29 which requires that the integrands converge in distribution to the relevant limit, and that their collection is uniformly integrable. Combination of the continuous mapping theorem and Lemmas 33, 12 and 30 yields convergence in distribution; application of the Hölder’s inequality, the polynomial bound on ζ\zeta, and Lemma 32 yields uniform integrability by Lemma 28, concluding the proof. ∎

which means that the expectation can be evaluated as a weighted sum of terms of the form

where a,b,c,d∈[ds]a,b,c,d\in[d^{s}], i,j,k,l∈LIi,j,k,l\in\mathcal{L}_{\mathcal{I}}, and s,t,u,v∈LXs,t,u,v\in\mathcal{L}_{\mathcal{X}}. We therefore only need to show convergence of these expectations. Substituting:

where εih\varepsilon_{i}^{h} are i.i.d. standard normal random variables, and re-purposing the i,ji,j indices, we have

Note that we can bound the integrands by a universal constant (Lemma 32), and thus we can focus only on the latter term on the r.h.s. We can thus turn to

and by Lemma 12 and the continuous mapping theorem

converges in distribution. By Lemma 30, this means that the integrand converges in distribution. Finally, to obtain the convergence of the expectation, we apply Theorem 29 where the required uniform integrability can be obtained by applying Hölder’s inequality and Lemma 32. ∎

The above formula suggests the desired result follows from Lemma 7:

Exchangeability requirement is satisfied by Lemma 13.

Zero mean and covariance follow from Lemma 14.

Convergence of variance is established in Corollary 15.

Under the assumptions of Theorem 3, φn,j\varphi_{n,j} are exchangeable over the jj index.

where we have w.l.o.g. assumed all matrices have been flattened as ⟨A,B⟩F=vect⁡(A)⊤vect⁡(B)\langle A,B\rangle_{F}=\operatorname{vect}(A)^{\top}\operatorname{vect}(B). The above could be further rewritten as a weighted sum of terms which take the following form:

Thanks to Lemma 33 and the continuous mapping theorem, we know that the integrand converges in probability to

which means we can conclude this proof by bounding this quantity by a constant independent of nn by Lemma 32. ∎

B.2 NTK convergence proof

Under the assumptions of Theorem 3 (including those stated at the beginning of Appendix B), for any a,b∈[ds]a,b\in[d^{s}], and x,x′∈Xx,x^{\prime}\in\mathcal{X}

Theorem 18 will be proven in the following two subsections.

The direct contribution of an attention layer can be expanded as

We prove convergence of each of these terms next.

Since dsd^{s} is fixed, we can focus on an arbitrary pair c1,c2∈[ds]c_{1},c_{2}\in[d^{s}]. Notice that by the continuous mapping theorem and Lemmas 33 and 30, the individual summands converge in distribution

Since dsd^{s} is fixed, we can focus on arbitrary c1,c2,d1,d2∈[ds]c_{1},c_{2},d_{1},d_{2}\in[d^{s}]. Rewriting the r.h.s. above for one such choice, we obtain

by Hölder’s inequality and exchangeability. Application of Lemma 32 allows us to bound the above r.h.s. by a constant independent of h,kh,k and nn as desired.

B.2.2 Indirect contribution

The indirect contribution of an attention layer can be expanded as

In the rest of this section, we drop the x,x′x,x^{\prime} from most of our equations so as to reduce the number of multi-line expressions. Continuing with the inner sum from above we obtain

which gives us four sums after multiplying out the terms inside the parenthesis, for each of which we prove convergence separately. Since the spatial dimension dsd^{s} does not change with nn, we will restrict our attention to an arbitrary fixed choice of a′,b′,c1,c2,d1,d2∈[ds]a^{\prime},b^{\prime},c_{1},c_{2},d_{1},d_{2}\in[d^{s}] throughout.

To make the notation more succinct, we define

Unlike in the proof of Lemma 22, the mean

only eliminates some of the sums. This issue can be resolved with the help of Lemma 24.

as desired. Application of Cheybshev’s inequality concludes the proof. ∎

With Sˉnh1h2\bar{S}_{n}^{h_{1}h_{2}} defined as in Lemma 24, we can revisit Equation 32

Note that the first two terms converge in probability to constant by Lemmas 33 and 24 and the continuous mapping theorem,

Note by the assumed continuity of ∇ζ\nabla\zeta, Theorem 3, Lemmas 33 and 24, the continuous mapping theorem, and Lemma 30, both the integrands converge in distribution, which, combined with the above derived bound and Theorem 29, implies

Analogously to the proof of Lemma 23, we define

Sˉnh⟶P0\bar{S}_{n}^{h}\overset{P}{\longrightarrow}0.

if τ=1\tau=1 (key and query weights are equal a.s.). Since each of the summands can be bounded by a constant indpendent of the i′,j′i^{\prime},j^{\prime} and nn indices by Lemma 32, we can restrict our focus to the terms for which i′≠j′i^{\prime}\neq j^{\prime}, yielding

To obtain convergence in probability, observe

Starting with the second expectation in Equation 36, we can use the assumed continuity of ∇ζ\nabla\zeta, Theorem 3, Equation 34, the continuous mapping theorem, and Lemma 30 to establish that the integrand converges in distribution to zero. Because ∇ζ\nabla\zeta is bounded by assumption, we can combine Hölder’s inequality and Lemma 32 to establish uniform integrability via Lemma 28 (see the proof of Lemma 26 for the bound on Sˉnh\bar{S}_{n}^{h}), and thus convergence of the expectation to zero by Theorem 29.

For the first expectation in Equation 36, note that the absolute value of the expectation can be upper bounded by

We provide a simple construction here, and expand on more realistic ones after the proof.

B.4 Auxiliary results

A sequence of real valued random variables (Xn)n≥1(X_{n})_{n\geq 1} is uniformly integrable if

If (Xn)n≥1(X_{n})_{n\geq 1} are uniformly integrable and Xn⇝XX_{n}\rightsquigarrow X, then XX is integrable and

meaning we can combine an argument analogous to the one above with Hölder’s inequality and exchangeability to obtain

Starting with the former, we can again replicate the argument from above, yielding

The obtain the convergence in probability, it is sufficient to show that

Appendix C Positional encodings

Similarly, the proof of Lemma 33 can be modified by observing that

C.2 NTK limit

Appendix D Residual attention