Vision Transformers provably learn spatial structure
Samy Jelassi, Michael E. Sander, Yuanzhi Li
Introduction
The empirical observation of 1(a) sets a central question: from a theoretical perspective, how do ViTs manage to learn these local connectivity patterns by simply minimizing their training loss using gradient descent from random initialization? While it is known that attention can express local operations as convolution (Cordonnier et al., 2019), it remains unclear how ViTs learn it. In this paper, we present a simple spatially-structured classification dataset for which it is sufficient (but not necessary) to learn the structure in order to generalize. We also present a simplified ViT model which we prove implicitly learns sparse spatial connectivity patterns when it minimizes its training loss via gradient descent (GD). We name this implicit bias patch association (defined in Definition 2.2). We prove that our ViT model leverages this bias to generalize. More precisely, we make the following contributions:
In Section 2, we formally define the concept of performing patch association, which refer to the ability of learning spatial connectivity patterns on a dataset.
In Section 3, we introduce a structured classification dataset and a simplified ViT model. This model is simplified in the sense that its attention matrix only depends on the positional encodings. We then present the learning problems we are interested in: empirical risk (realistic setting) and population risk (idealized setting) minimization for binary classification.
In Section 4, we prove that a one-layer single-head ViT model trained with gradient descent on our synthetic dataset performs patch association and generalizes, in the idealized (Theorem 4.1) and realistic (Theorem 4.2) settings. We present a detailed proof, based on invariance and symmetries of coefficients in the attention matrix throughout the learning process.
In Section 5, we show (Theorem 5.1) that after pre-training in our synthetic dataset, our model can be sample-efficiently fine-tuned to transfer to a downstream dataset that shares the same structure as the source dataset (and may have different features).
On the experimental side, we validate in Section 6 that ViTs learn spatial structure in images from the CIFAR-100 dataset, even when the pixels of the images are permuted. This result validates that, in contrast to CNNs, ViTs learn a more general form of spatial structure that is not limited to local patterns (Figure 3). We finally show that our ViT model –where the attention matrix only depends on the positional encodings– is competitive with the vanilla ViT on the ImageNet, CIFAR-10/100 and SVHNs datasets (Section 6 and Section 6).
Related work
Many computer vision architectures can be considered as a form of hybridization between Transformers and CNNs. For example, DeTR (Carion et al., 2020) use a CNN to generate features that are fed to a Transformer. (d’Ascoli et al., 2021) show that self-attention can be initialized or regularized to behave like a convolution and (Dai et al., 2021; Guo et al., 2021) add convolution operations to Transformers. Conversely, (Bello et al., 2019; Ramachandran et al., 2019; Bello, 2021) introduce self-attention or attention-like operations to supplement or replace convolution in ResNet-like models. In contrast, our paper does not consider any form of hybridization with CNN, but rather a simplification of the original ViT to explain how ViTs learn spatially structured patterns using GD.
A long line of work consists in analyzing the properties of ViTs, such as robustness (Bhojanapalli et al., 2021; Paul and Chen, 2021; Naseer et al., 2021) or the effect of self-supervision (Caron et al., 2021; Chen et al., 2021b). Closer to our work, some papers investigate why ViTs perform so well. Raghu et al. (2021) compare the representations of ViTs and CNNs and Melas-Kyriazi (2021); Trockman and Kolter (2022) argue that the patch embeddings could explain the performance of ViTs. We empirically show in Section 6 that applying the attention matrices to the positional encodings – which contains the structure of the dataset – approximately recovers the baselines. Hence, our work rather suggests that the structural learning performed by the attention matrices may explain the success of ViTs.
Early theoretical works have focused on the expressivity of attention. (Vuckovic et al., 2020; Edelman et al., 2021) addressed this question in the context of self-attention blocks and (Dehghani et al., 2018; Wei et al., 2021; Hron et al., 2020) for Transformers. On the optimization side, (Zhang et al., 2020) investigate the role of adaptive methods in attention models and (Snell et al., 2021) analyze the dynamics of a single-head attention head to approximate the learning of a Seq2Seq architecture. In our work, we also consider a single-head ViT trained with gradient descent and exhibit a setting where it provably learns convolution-like patterns and generalizes.
The question we address concerns algorithmic regularization which characterizes the generalization of an optimization algorithm when multiple global solutions exist in over-parametrized models. This regularization arises in deep learning mainly due to the non-convexity of the objective function. Indeed, this latter potentially creates multiple global minima scattered in the space that vastly differ in terms of generalization. Algorithmic regularization appears in binary classification (Soudry et al., 2018; Lyu and Li, 2019; Chizat and Bach, 2020), matrix factorization (Gunasekar et al., 2018; Arora et al., 2019), convolutional neural networks (Gunasekar et al., 2018; Jagadeesan et al., 2022), generative adversarial networks (Allen-Zhu and Li, 2021), contrastive learning (Wen and Li, 2021) and mixture of experts (Chen et al., 2022). Algorithmic regularization is induced by and depends on many factors such as learning rate and batch size (Goyal et al., 2017; Hoffer et al., 2017; Keskar et al., 2016; Smith et al., 2018; Li et al., 2019), initialization Allen-Zhu and Li (2020), momentum (Jelassi and Li, 2022), adaptive step-size (Kingma and Ba, 2014; Neyshabur et al., 2015; Daniely, 2017; Wilson et al., 2017; Zou et al., 2021; Jelassi et al., 2022), batch normalization (Arora et al., 2018; Hoffer et al., 2019; Ioffe and Szegedy, 2015) and dropout (Srivastava et al., 2014; Wei et al., 2020). However, all these works consider the case of feed-forward neural networks which does not apply to ViTs.
Defining patch association
The goal of this section is to formalize the way ViTs learn sparse spatial connectivity patterns. We thus introduce the concept of performing patch association for a spatially structured dataset.
Setting to learn patch association
In this section, we introduce our theoretical setting to analyze how ViTs learn patch association. We first define our binary classification dataset and finally present the ViT model we use to classify it.
We sketch a data-point of in Section 3. Our dataset can be viewed as an extreme simplification of real-world image datasets where there is a set of adjacent patches that contain a useful feature (e.g. the nose of a dog) and many patches that have uninformative or spurious features e.g. the background of the image. We make the following assumption on the parameters of the data distribution.
Assumption 2 may be justified by considering a "ViT-base-patch16-224" model Dosovitskiy et al. (2020) on ImageNet. In this case, , . is set to have . is chosen so that there are more spurious features than informative ones (low signal-to-noise regime) which makes the data non-linearly separable. Our dataset is non-trivial to learn since generalized linear networks fail to generalize, as shown in the next theorem (see Appendix J for a proof).
We now define our simplified ViT model for which we show in Section 4 that it implicitly learns patch association via minimizing its training objective. We first remind the self-attention mechanism that is ubiquitously used in transformers.
the sum of patches and positional encodings i.e.
In this paper, our ViT model relies on a different attention mechanism –the "positional attention"– that we define as follows.
Positional attention isolates positional encoding from data : encodes the dynamics of and tracks whether patch association is learned. encodes the data-dependent part and monitors whether the feature is learned. Indeed, given its highly non-linear nature with respect to the input, directly analyzing self-attention is difficult. Yet, positional attention is similar to self-attention. As this latter, positional attention is also permutation-invariant and processes all tokens simultaneously. Besides, positional attention also computes a score matrix between the different tokens. This similarity matrix is also normalized in a sparse manner with the Softmax operator. The only aspect that positional attention misses from self-attention is the fact that does not depend on the input. Nevertheless, we empirically show that our positional attention model competes with self-attention in Section 6. Lastly, we make the following simplification in the parameters to ease our analysis.
In Simplification 3.1, we set and to the identity so that This Gram matrix encodes the spatial patterns learned by the ViT as shown in 1(a). Besides, since fitting the labeling function requires to learn one feature , it is sufficient to parameterize with a vector . Also, although and is trainable, we choose for simplicity to only optimize over Besides, we leave the ’s fixed because Softmax is invariant under the uniform shift of the input. Under Simplification 3.1, our simplified ViT model is then a two attention layer with a single head:
Given a dataset sampled from , we solve the empirical risk minimization problem for the logistic loss defined by:
Instead of directly analyzing (E), we introduce a proxy where we minimize the population risk
We refer to (E) as the realistic problem while (P) as the idealized problem.
We solve (P) and (E) using gradient descent (GD) for iterations. The update rule in the case of (P) for and is
where is the learning rate. A similar update may be written for (E). We now detail how to set the parameters in (GD).
Idealized case: where and for
Learning spatial structure via matching the labeling function
As announced above, we show that our ViT (T) implicitly learns patch association and fits the labeling function by minimizing the training objective. We first study the dynamics in (P). Using the analysis in the idealized case, we then characterize the solution found in the realistic problem (E).
In this section, we analyze the dynamics of (P). Our main result is that after minimizing (P), our model (T) performs patch association while generalizing.
Assume that we run GD on (P) for iterations with parameters set as in Parametrization 3.1. With high probability, the ViT model (T)
We now sketch the main ideas to prove the theorem for which one can refer to Appendix D for a complete proof.
In (P), we take the expectation over . Since (T) is permutation-invariant and the data distribution is symmetric, we can thus dramatically simplify the variables in (P). An illustration of this is the next lemma that shows that can be reduced to three variables in (P).
for all ,
In summary, Lemma 4.1 and Lemma 4.2 imply that instead of optimizing over and in (P), we can instead consider the scalar variables , and . The remaining of this section consists in analyzing the dynamics of these three quantities.
We first analyze the dynamics of and . To this end, we introduce the following terms:
Let . The attention weights and satisfy:
Lemma 4.3 shows that the increment of is larger than the one of . Since , this implies that for all This observation proves the first item of Theorem 4.1. We now explain how learning patch association leads to highly correlated with
Event I: At the beginning of the process, the update of is larger than the one of which implies that only updates during this first phase. We show that increases until a time where it reaches some threshold (Lemma D.2). At this point, the model is nothing else than a generalized linear model that would not generalize because there are much more noisy tokens than signal ones (see Theorem 3.1).
Event II: During this phase, the attention weights must update. Indeed, assume by contradiction that the stay around initialization and that is optimal i.e. where Then, the predictor we would have is
Event III: Because we have , we again have as in Phase I (Lemma D.11). Thus, increases again until the population risk becomes a .
Our mechanism highlights two important aspects that are proper to attention models:
because of the initialization and the data structure, we have patch association for any time (Lemma 4.3).
our ViT model uses patch association to minimize the population loss (Event III). Without patch association, the model would only be a generalized linear model that does not minimize the loss.
2 From the idealized to the realistic learning process
The real learning process differs from the idealized one in that we have a finite number of samples and we initialize both and as Gaussian random variables. Using a polynomial number of samples, we show that (T) still learns patch association and generalizes.
Similarly to Li et al. (2020), the proof introduces a "semi-realistic" learning process that is a mid-point between the idealized and realistic processes. We show that and are close to their semi-realistic counterparts – see Appendix E for a complete proof. Figure 2 numerically illustrates Theorem 4.2.
Patch association yields sample-efficient fine-tuning with ViTs
A fundamental byproduct of our theory is that after pre-training on a dataset sampled from , our model (T) sample-efficiently transfers to datasets that are structured as but differ in their features.
Let a downstream data distribution defined as in Assumption 1 such that its underlying feature is with and potentially different from . In other words, the downstream and source distributions share the same structure but not necessarily the same feature. We sample a downstream dataset from .
We consider the model (T) pre-trained as in subsection 4.2. We assume that is kept fixed from the pre-trained model and we only optimize the value vector to solve:
The proofs of Theorem 5.1 and Theorem 5.2 are in Appendix F. These theorems hightlight that learning patch association is required for efficient transfer. We believe that they offer a new perspective on explaining why ViTs are widely used in transferring to downstream tasks. While it is possible that ViTs learn shared (with the downstream dataset) features during pretraining, our theory hints that learning the inductive bias of the labeling function is also central for transfer.
Numerical experiments
In this section, we first empirically verify that ViTs learn patch association while miniziming their training loss. We then numerically show that the positional attention mechanism competes with the vanilla one on small-scale datasets such as CIFAR-10/100 (Krizhevsky et al., 2009), SVHN (Netzer et al., 2011) and large-scale ones such as ILSVRC-2012 ImageNet (Deng et al., 2009). For the small datasets, we use a ViT with 7 layers, 12 heads and hidden/MLP dimension 384. For ImageNet, we train a "ViT-tiny-patch16-224" Dosovitskiy et al. (2020). Both models are trained with standard augmentations techniques (Cubuk et al., 2018) and using AdamW with a cosine learning rate scheduler. We run all the experiments for 300 epochs, with batch size 1024 for Imagenet and 128 otherwise and average our results over 5 seeds. We refer to Appendix A for the training details.
Test accuracy obtained with a ViT using vanilla attention (ViT) and positional attention (Ours) on CIFAR-10 (1), CIFAR-100 (2) and SVHN (3). Our model competes with the vanilla ViT. Patch size 4 and average over 10 seeds for this experiment. We numerically verify that ViTs using positional attention compete with those with vanilla attention. In Section 3, we introduced positional attention to define our theoretical learner model. Section 6 and Section 6 show that ViTs using positional attention compete with vanilla ViTs on a range of datasets. These experiments strengthen our intuition that for images, having an attention matrix that only depends on the positional encodings is sufficient to have a good test accuracy.
Conclusion, limitations and future works
Our work is a first step towards understanding how Transformers learn tailored inductive biases when trained with gradient descent. Our analysis heavily relies on the positional attention mechanism that disentangles patches and positional encodings. In practice, self-attention mixes these two quantities. An interesting direction is to understand the impact of patch embeddings on the inductive bias learned by ViTs. Moreover, our experiment on the Gaussian data shows that ViTs do not always learn the correct inductive bias under Definition 2.1: characterizing the distributions under which ViTs recover the structure of the function is an important question. Lastly, this work also paves the way to many extensions beyond convolution. For example, can ViTs learn other inductive biases? What are the inductive biases learnt by Transformers in NLP? Answering those questions is central to better understand the underlying mechanism of attention.
Acknowledgments and Disclosure of Funding
The authors would like to thank Boris Hanin for helpful discussions and feedback on this work.
References
Checklist
Do the main claims made in the abstract and introduction accurately reflect the paper’s contributions and scope? [Yes] See Section 4 and Section 5.
Did you describe the limitations of your work? [Yes] See Conclusion, limitations and future works.
Did you discuss any potential negative societal impacts of your work? [N/A] This is a theory paper.
Have you read the ethics review guidelines and ensured that your paper conforms to them? [Yes]
If you are including theoretical results…
Did you state the full set of assumptions of all theoretical results? [Yes] See Section 3.
Did you include complete proofs of all theoretical results? [Yes] See Appendix.
Did you include the code, data, and instructions needed to reproduce the main experimental results (either in the supplemental material or as a URL)? [Yes] See supplementary material.
Did you specify all the training details (e.g., data splits, hyperparameters, how they were chosen)? [Yes] See Appendix.
Did you report error bars (e.g., with respect to the random seed after running experiments multiple times)? [Yes] See Section 6
Did you include the total amount of compute and the type of resources used (e.g., type of GPUs, internal cluster, or cloud provider)? [Yes] See Appendix.
If you are using existing assets (e.g., code, data, models) or curating/releasing new assets…
If your work uses existing assets, did you cite the creators? [Yes]
Did you mention the license of the assets? [N/A]
Did you include any new assets either in the supplemental material or as a URL? [N/A]
Did you discuss whether and how consent was obtained from people whose data you’re using/curating? [N/A]
Did you discuss whether the data you are using/curating contains personally identifiable information or offensive content? [N/A]
If you used crowdsourcing or conducted research with human subjects…
Did you include the full text of instructions given to participants and screenshots, if applicable? [N/A]
Did you describe any potential participant risks, with links to Institutional Review Board (IRB) approvals, if applicable? [N/A]
Did you include the estimated hourly wage paid to participants and the total amount spent on participant compensation? [N/A]
Appendix A Additional experimental details
In this section, we provide additional details on our experiments and additional plots.
We used Pytorch and Nvidia Tesla V100 GPUs. We conduct experiments on small-scale (CIFAR-10/100 and SVHN) and large-scale datasets (ImageNet). The choice of architecture and training parameters depend on the size of the dataset as we detail below.
We use the code available at https://github.com/omihub777/ViT-CIFAR. The model is made of 7 layers, 12 heads, hidden and MLP dimension 384, dropout 0. We use "mean-pooling" and not the CLS pooling. We set the patch size to 2 in the experiment Figure 3 and to 4 in the experiment Section 6. Indeed, we empirically found that setting patch size 4 was the optimal choice. We apply label smoothing [Szegedy et al., 2016] with coefficient 0.1 and do not apply any cutmix [Zhang et al., 2017] nor mixup [Yun et al., 2019]. We use Adam [Kingma and Ba, 2014] as optimizer and set the learning rate to , minimum learning rate to , to , to , batch size to , weight decay to , number of warmup epochs to 5 and number of total epochs to 200. The scheduler is a cosine learning rate. We used the AutoAugment procedure [Cubuk et al., 2018] as in the repository to generate data augmentations. The model has been trained over a single GPU.
Regarding the convolutional models in the experiment Figure 3, we trained a ResNet-18 and a VGG-19 with batch normalization. We trained the two architectures using the same training procedure and hyperparameters as for the ViT.
A.2 Additional plots
In Figure 3, we plot the positional encoding similarities for a few patches. Figure 4 provides these plots for all the patches. One should think of Figure 3 as a Figure displaying just two of the arrays present in Figure 4. We consistently verify that the ViT is always able to recover the convolution-like patterns which shows that it is able to learn the right patch association.
Appendix B Induction hypothesis
In this section, we present the induction hypothesis that we use in the analysis of the idealized case. This hypothesis is ultimately proved in subsection D.7.
During the idealized learning process, the following holds for .
the sofmax denominator is large i.e.
is not too small i.e.
and are in a good range i.e.
Appendix C Notations
In this section, we introduce the different notations used in the proofs.
We first define notations that are used everywhere in the appendix.
Loss for a data-point :
We now provide notations used in the analysis of the idealized case.
for ,
We now provide notations used in the analysis of the realistic case.
Given a data-point and ,
Appendix D Learning process in the idealized setting
We divide the idealized learning process as follows.
Event I (, subsection D.2): at initialization, is small. Therefore, the sigmoid is large. Besides, around , it stays constant i.e. in (GD-). This implies that which yields to increase until reaching a specific value where the sigmoid is not constant anymore.
Event II (, subsection D.3): at time , is large. This fact along with Lemma 4.3 imply that increases. Eventually, becomes large enough so that .
Event III (, subsection D.4): Since , increases again. It increases until the population risk is at most
After iterations, is large and the population risk thus converges (subsection D.5). Since the logistic loss is a surrogate for the 0-1 loss, we prove that the learner model fits the labeling function (subsection D.6) which implies the first statement of Theorem 4.1.
Remark : Since we initialize , Lemma D.7 implies that we can overlook the linear part of the activation in this section. Therefore, we only consider in the idealized process.
A first question that arises is: starting from , what is the value of that makes the sigmoid non-constant? The following lemma addresses this question.
The value at which the sigmoid becomes non-constant is:
where Since is -Lipschitz, we rewrite (3) as:
Let . For all , we have Therefore, is updated as
Consequently, is non-decreasing and after iterations, we have for
For , we know that the sigmoid is constant. We apply Lemma D.4 and Lemma D.6 to respectively bound and in the update of .
(8) indicates that is a non-decreasing sequence. Therefore, there exists a time such that . Using Lemma K.1, the time is equal to:
In this section, we present the auxiliary lemmas needed to prove the main results of subsection D.2. We first present a lemma that bounds the learner model.
Let . The learner model is bounded for all as:
We successively apply Lemma D.4, Induction Hypothesis B.1 and Lemma D.5 to bound (10).
Finally, we apply Induction Hypothesis B.1 in (11) to obtain the desired result. ∎
We now present lemmas that bound and .
We now sum (LABEL:eq:jcneiencae) and apply Lemma K.3 to obtain:
We finally apply Induction Hypothesis B.1 to have in (15) and get
We finally plug (13) and (16) in (LABEL:eq:jwfejfw) and obtain:
Let We have
By definition of , we have:
Let and . Assume that . Then, we have:
We remind that the derivative of the activation function . We first remark that for all . Besides, we have:
In this section, we show the increase of for leads to the increase of . At time is significantly large.
Let and Using Corollary G.1 and Induction Hypothesis B.1, satisfies:
Summing (22) for yields
We successively apply Lemma D.9 and for to lower bound (23) to obtain:
We apply Induction Hypothesis B.1 in (24) to obtain a bound on .
(25) shows that is an non-decreasing sequence. We thus deduce the time such that
We now prove the second part of the lemma. We respectively apply Lemma D.4 and Lemma D.5 to bound and .
(27) implies for all ,
In this section, we present the auxiliary lemmas needed to prove the main results in subsection D.3.
Let Then, we always have
For , increases and eventually satisfies (Lemma D.2). However, for , may be non-increasing. Here, we want to quantify the maximum amount of decrease for . The worst-case scenario is when . We bound by using Lemma D.4, Lemma D.10 and Lemma D.5.
We now apply Induction Hypothesis B.1 in (28) and get:
At time we potentially have . In this case, starts to increase again because it is in the range of ’s that satisfies Event I (and therefore the update rule in Lemma D.2 holds). Thus, for all we have ∎
Let be a data-point. We distinguish two cases:
: we apply Lemma K.6 which implies Since the population loss is , this implies the aimed result.
: we have necessarily since the sigmoid function is large for non-positive values.
For , increases until reaching In this section, we show that this implies that increases again.
Since (Lemma D.8), the update of is:
We say that the sigmoid term is small for a constant that satisfies
Intuitively, (32) means that the sum of the sigmoid terms for all time steps is bounded (up to a logarithmic dependence). In our case, by using (32) and Lemma D.3, the sigmoid is small when
D.5 Convergence rate of the population loss
Let . Then, the population loss linearly converges to zero i.e.
Let’s now assume by contradiction that for , we have:
For , we know that is non-decreasing which implies that is also non-decreasing. Since is non-increasing, this implies for that
Plugging (40) in the update (36) yields for :
Let . We now sum (41) for and obtain:
where we used (Lemma D.8) and (42) in the last inequality. We now apply Lemma K.6 and obtain:
Given the values of , we finally have:
which contradicts (39). Therefore, we obtain the convergence rate:
We apply Lemma D.14 to bound the left-hand side of (46) and get the aimed result. ∎
The proof is similar to the one of Lemma D.3. We apply Lemma D.4, Lemma D.8 and Lemma D.5 and get:
where we applied Lemma K.7 in the last inequality. ∎
D.6 Fitting the labeling function
We now show that the learner model fits the labeling function.
Since the logistic loss is a surrogate for the 0-1 loss, we have:
We now apply Lemma D.13 to bound the right-hand side of (49). Given the value of , we have:
D.7 Proof of the induction hypothesis
In this section, we prove Induction Hypothesis B.1.
We start by proving that for all Let and Using Corollary G.2 and Induction Hypothesis B.1, we upper bound as:
Summing (52) for and using lead to
We now apply Lemma D.16 to bound the sum of ’s in (53).
Given the values of the different parameters, (54) implies that
We now prove Since is non-decreasing (Corollary G.1), we have for all We now prove the upper bound on . We assume that for all , Let’s show this inequality for . Using Corollary G.1, we have:
We now apply the induction hypothesis in (55) and get:
where we used the inequality for in (57). Given the values of the different parameters, we deduce that
Lastly, we prove that for Since , and , we have:
The sum of the ’s is bounded as:
We first decompose the sum of ’s.
We apply Lemma D.2 and Lemma D.8 to rewrite (60).
Now, we aim to obtain the value of the last summand in (61). Using Corollary G.1, we have
We finally apply Lemma D.8 and Induction Hypothesis B.1 in (62) to get:
Appendix E From idealized to real learning process
Since we randomly initialize with tiny variance, we need to take into account the linear part of the activation function. Lemma E.8 shows that we can overlook the power part of the activation and consider as long as .
Consequently, is non-decreasing and after iterations, we have for
Let . We apply Lemma E.5 and Lemma E.6 to respectively bound and in the update of .
(64) indicates that is a non-decreasing sequence. Therefore, there exists a time such that . Summing (64) for yields \mathscr{T}=\Theta\Big{(}\frac{1}{\eta\nu^{(p-2)/(p-1)}e^{\beta}}\Big{)}. ∎
We now show that for , the orthogonal component stays small.
Assume that we run GD on the empirical risk (E) for iterations with parameters set as in Parametrization 3.1. For , the orthogonal component satisfies
Let and . The projected update of satisfies:
We use the 1-Lipschitzness of the sigmoid function and get:
where we applied Induction Hypothesis B.1 in (69). Since with high probability, , \big{|}\langle\bm{u}^{(t)},\sum_{r=1}^{D}\bm{\xi}_{r}\rangle\big{|}\leq\sqrt{D\log(d)}\sigma, we finally have:
Combining the bounds on Summands 1 and 2 yields the aimed result. ∎
We now use Lemma E.2 to show that stays small.
Unraveling Lemma E.2 for and using (with high probability) leads to:
where we used for and Plugging the value of in (70) yields the aimed result. ∎
We finally show that remains tiny for
Let . We have
We remind that the GD update of is
The proof is by induction. We assume that We first apply Cauchy-Schwarz on (71) and get:
Using the induction hypothesis, we have Thus, we have
We now use Lemma E.1 and Lemma E.3 in (74) and get:
E.1.4 Auxiliary lemmas
This implies for all
We successively apply Lemma E.4 and Lemma K.3 to get:
Let . The sum of ’s is bounded as:
Let We sum the update rule of (Lemma E.1) and obtain: Summing again this update yields the aimed result.
Let and . Assume that . Then, we have:
We remind that the derivative of the activation function . We first remark that for all . Besides, we have since is even,
E.2 Coupling between the semi-idealized and realistic processes (t∈[𝒯,T]𝑡𝒯𝑇t\in[\mathscr{T},T])
In this section, we aim to bound the realistic iterates and for For this reason, we introduce a "semi-idealized" learning process (subsubsection E.2.1) which may be viewed as a mid-point between the idealized and realistic process. We first bound the iterates in this process. Then, using this process, we show that (subsubsection E.2.2) and (subsubsection E.2.4) stay small. Here, is the semi-idealized attention matrix coefficient. Finally, since and are small, the final iterates and are equal (subsubsection E.2.6) and thus, the model fits the labeling function (subsubsection E.2.7).
We define an intermediate learning process that we refer to as the "semi-idealized" process. This process starts at time involves two parameters: the semi-idealized value vector and semi-idealized attention matrix defined as
the value vector is fixed and satisfies for
is a trainable parameter and is initialized as for .
Therefore, the only trainable parameter in this process is . In the semi-idealized process, we minimize the population risk
We remark that such process present similarities to the idealized case. In particular, it satisfies all the invariance and symmetry properties from Lemma 4.1. We thus define
Therefore, and are respectively updated as in Lemma G.1 and Lemma G.2. We define also the softmax terms
We finally assume Induction Hypothesis B.1 for this process. This latter can be proved using the same arguments as in subsection D.7.
We previously showed in Lemma E.3 that is small in the initial steps. We now show that it stays small during the whole process.
We now proceed to the proof of Lemma E.9. We first characterize the recursion satisfied by
Assume that we run GD on the empirical risk (E) for iterations with parameters set as in Parametrization 3.1. Then, satisfies for
Let , and The projected update of satisfies:
With high probability, , . We thus get:
We apply Lemma E.12 to bound the local change of the sigmoid in (89) which yields:
We apply Induction Hypothesis B.1 to bound the softmax terms in (90). Besides, with high probability, we have , . Thus, we have:
We combine the bounds on the three summands to obtain the recursion of . ∎
We now prove Lemma E.11 that gives the final bound on for
We bound in the following two regimes: and
Lemma E.20 provides the update of during this time phase. We thus apply Lemma K.2 to bound the product term in (92).
Plugging (93) in (92) yields a bound on
E.2.3 Auxiliary lemmas
In this section, we prove the Lipschitzness of the function appearing in the proof of Lemma E.10.
Let The derivative of with respect to a variable is:
(97) implies a bound on . Indeed, since , we have:
(98) shows that is -Lipschitz. ∎
Here, we bound the gap in attention coefficients between the realistic and semi-idealized cases.
We now detail the steps to prove Lemma E.13. We first provide the recursion that satisfies.
Assume that we run GD on the empirical risk (E) for iterations with parameters set as in Parametrization 3.1. Then, the discrepancy satisfies for ,
In this proof, we maintain the hypothesis that is small. We will eventually prove this statement in Lemma E.15. Let such that . Using GD, satisfies:
Since is small (Lemma E.11), we can show that:
where we used in (115) and Induction Hypothesis B.1 in (116). Using Lipschitz inequalities, we can further expand (116) as a function of the coefficients from and . However, is small and we only want terms of order 1 in in (116). Therefore, the only term of order 1 that remains is:
Bounding the expectation in (LABEL:eq:wefcerffer) as in the proof of Lemma G.1 yields . We now bound . We therefore apply (Lemma E.16) and get:
where we used in (120). We can further expand (120), keep the terms of first order in and get
We now bound . Using the Lipschitz property of the softmax (Lemma E.17), we have:
where we used in (121). Using the same arguments as above, we obtain
The bound on can be derived as above. We again use the Lipschitz property of softmax (Lemma E.17) which leads to
Plugging the bounds on Summands 1, 2 and 3 in the original decomposition of \big{|}\widehat{A}_{a,b}^{(t+1)}-\widecheck{A}_{a,b}^{(t+1-\mathscr{T})}\big{|} yields the bound on . The second part of the lemma is obtained using Lemma E.15.
Let such that for – we proved the existence of in Lemma E.11. We bound when and
We now apply Induction Hypothesis B.1 to simplify (LABEL:eq:jfejfe) and get:
We then apply Lemma K.2 to bound the product term in (125). We obtain:
E.2.5 Auxiliary lemmas
We finally apply the generalized mediant inequality in (130) and get:
Let . The difference of softmax is bounded as:
where we used the mediant inequality in the last inequality of (131). Since the exponential function is non-decreasing, we deduce:
Analog of Event I (Lemma D.2): for .
Analog of Event III (Lemma D.11): for .
These three lemmas imply that at time , the realistic iterates are very close to the ideal ones. Therefore, they incur nearby test loss and thus the realistic model generalizes. We now proceed to the proof of
In order to analyze the dynamics of , we first show that the gradient (with respect to ) in the realistic learning process is very close to the one in the semi-idealized one.
Let . With high probability, we have
Lemma E.19 shows that we can use the gradient from the semi-idealized process to analyze the dynamics of in the real process. Therefore, we can derive similar updates for as in Lemma D.2, Lemma D.8 and Lemma D.11.
Let . Therefore, is updated as
Consequently, is non-decreasing and after iterations, we have for
The following lemma is useful to prove Lemma E.10 and Lemma E.14.
The result is obtained by applying Lemma K.1 to (142). We have:
E.2.7 The realistic model fits the labeling function
In the realistic case, the model fits the labeling function i.e.
We bound the population risk . We have:
We can further expand (146) as in the proof of Lemma D.15 and deduce the aimed result. ∎
To prove Lemma E.24, we use the following auxiliary lemma.
After iterations, the population risk in the semi-idealized case converges i.e.
The proof is similar to the one of Lemma D.15. ∎
For all sampled from , we have
Since is Lipschitz on a bounded domain, we have:
since . Using Cauchy-Schwarz inequality, (149) simplifies as:
We again use the Lipschitzness of the power function and get:
Appendix F Transfer Learning
In this section, we show that a transformer that has been pre-trained on a structured dataset require a few samples to generalize in a new dataset sharing the same structure.
Actually, even one step of the update using normalized gradient descent on can already achieve test accuracy We know that for a datum , the gradient of with respect to is
Since we have and Thus, the gradient (153) simplifies to
By symmetry of the , we know that
Now, by standard concentration inequality, we know that for i.i.d. samples , with high probability
where comes from the feature noise
and comes from the noise:
Therefore, if we update using normalized GD:
We can prove the test accuracy is small using the same proof as in Lemma D.15, where we show that:
: Sample where each i.i.d. w.p. , w.p. and otherwise.
: Sample a set uniformly at random from of size , set all for , and sample other i.i.d. w.p. , w.p. and otherwise.
We can easily see that as long as , then
Appendix G Gradient descent updates in the idealized process
In this section, we derive the gradient descent updates of in the idealized learning process.
Let be the time where the population loss is at most and Then, satisfies the update
Since , we rewrite (164) as:
Regarding the sum inside , we use Lemma G.3 which shows:
By using (165) and (166), we finally obtain:
We now bound the sum inside . This sum is actually equal to the outside sum and we can therefore use the bound (168). Therefore, the overall gradient is bounded as:
We lastly apply Induction Hypothesis B.1 to show that (170) is less or equal to We now bound the sum inside .
We now bound the sum inside .
We lastly apply Lemma G.3 to show that (LABEL:eq:jewnwcdw) is bounded by Therefore, the overall gradient is bounded as
We now bound the derivative of the population loss. Using Tower property and Lemma I.1, we have:
Therefore, the derivative of the loss in is:
Let be the time where the population loss is and Let . The update of satisfies:
Lemma G.1 provides the update rule of .
Let be the time where the population loss is at most and Then, satisfies the update
We now bound the sum inside . This sum is actually equal to the outside sum and we can therefore use the bound (179). Therefore, the overall gradient is bounded as:
We now bound the sum inside . We successively apply Lemma K.3, triangle inequality and Lemma G.3 to obtain:
Thus, we use (181) and (182) to obtain a bound on the derivative.
We now bound the sum inside .
Using (184) and (185), we obtain a bound on the derivative.
We now bound the sum inside . We apply Lemma G.3 to obtain:
We combine (187) and (LABEL:eq:fjqjeq) and obtain:
We now bound the sum inside . This sum is actually equal to the outside sum outside and we can therefore use the bound (LABEL:eq:frekfwew). Thus, the derivative is bounded as:
We now bound the sum inside the power term. This sum is actually equal to the sum outside the power term and we can therefore use the bound (LABEL:eq:redeww). Thus, the derivative is bounded as:
We now bound the sum inside .
Using (LABEL:eq:jfnrorjw3) and (195), the bound on the derivative is:
We now bound the sum inside the power term. We apply Lemma G.3 and get:
We plug (LABEL:eq:fjeiwbf) and (198) to obtain the derivative.
We now bound the derivative of the population loss. Using Tower property and and Lemma I.1, we have:
Since events a and e are the ones with highest probabilities, the derivative of the loss is bounded by the expectations conditioned on these events. We have:
We now apply Induction Hypothesis B.1 and Lemma K.3 and finally obtain:
Let be the time where the population loss is and The update of satisfies:
Lastly, we apply Induction Hypothesis B.1 to replace by its value in (203) and thus obtain the aimed result. ∎
G.3 Auxiliary lemmas
Let In the idealized learning process, with high probability, we have:
We first bound the sum with factor . Using Induction Hypothesis B.1 and Lemma K.3, we have:
We now bound . Using Induction Hypothesis B.1, we have:
Let . In the idealized learning process, we have with high probability:
We successively apply Lemma K.3 and Induction Hypothesis B.1 to get the desired bound. Indeed, we have:
Appendix H Gradients
In this section, we present the gradients of the loss with respect to and .
Let be a data-point. Then, the gradient of with respect to is:
Let be a data-point and . The derivative of with respect to is:
Appendix I Invariance of the problem
Using Induction Hypothesis B.1, we simplify (LABEL:eq:fewjoejdwe) as
Therefore, (LABEL:eq:fe) implies that thus proving the induction hypothesis.
I.2 Invariance by permutation
Let and be two permutations and . Let . Then, we have:
permutation-invariant distribution: has the same distribution as
permutation-invariant model:
Let be a data-point. Using Lemma 4.1, we rewrite as
Appendix J Justification of our data distribution
In this section, we justify why the distribution (Assumption 1) is relevant. We first show that linear classifiers poorly generalize (subsection J.1). We then show that there exists classifiers that generalize without learning patch association (subsection J.2).
For every data point , consider , it is very easy to see that for every integer , as long as , we have that:
Consider two independently sampled data points, with label respectively, consider the event when and all the noises of satisfies , then we know that
By Eq (214) we also know that the density of and under the data-generation distribution satisfies
J.2 Classifiers fitting the labelling function without patch association
Appendix K Technical lemmas
In this section, we present the technical lemmas used in the paper.
Let be a positive sequence defined by the following recursions
where is the initialization, is an integer and . Let such that Then, the time such that for all is:
We use the fact that in (223) and obtain:
Now, we want to bound . Using again the recursion and , we have:
Combining (224) and (225), we get a bound on
Now, let’s find a bound for . Starting from the recursion and using the fact that for we have:
On the other hand, by using we upper bound as follows.
Besides, we know that . Therefore, we upper bound as
We now sum (230) for , use (226) and obtain:
Lastly, we know that satisfies which implies in (231). ∎
Let be a positive sequence defined by the following recursions
where , is an integer and . Let such that and be the time such that for all . Assume that . Then, we have for :
Since is a non-decreasing sequence, (232) satisfies:
We now sum (233) for and get:
Since , we have \big{(}1+m(z^{(t)})^{k-1}\big{)}^{\frac{A(z^{(t)})^{\kappa}}{m}}\geq 1+A(z^{(t)})^{k-1+\kappa}. We thus lower bound (234) as:
On the other hand, by using and , we have the following upper bound.
Since , (236) is finally bounded as:
We combine (235) and (LABEL:eq:ewjeof) to obtain:
We replace by and by in (239) to get the aimed result. ∎
K.2 Probabilistic lemmas
K.3 Logarithmic inequalities
We upper bound (242) by applying :
We obtain the final bound by applying Lemma K.6 to (243).
We lower bound (242) by using :
We obtain the final bound by applying Lemma K.6 to (244). ∎
Let Assume that Then, we have: