How Do Transformers Learn Topic Structure: Towards a Mechanistic Understanding

Yuchen Li, Yuanzhi Li, Andrej Risteski

INTRODUCTION

The transformer architecture (Vaswani et al., 2017) is a critical building block of many leading approaches to natural language processing (Devlin et al., 2019; Brown et al., 2020), and other domains such as vision (Dosovitskiy et al., 2021) and protein structure prediction (Jumper et al., 2021). While the NLP community has produced a large body of work on probing and visualizing trained networks (Hewitt & Manning, 2019; Clark et al., 2019; Tenney et al., 2019; Kovaleva et al., 2019), we still have little formal understanding of the mechanisms by which transformers, trained with simple gradient-descent based algorithms, learn from their training data. The challenge is that the training dynamics are non-trivial, even for relatively simple structured data distributions, and even for simple (e.g. 1-layer) transformers.

In particular, we study semantic structure, as understood through the lens of co-occurrences of words, and their topical structure. Precisely, if we fit topics to a real-life corpus like Wikipedia using a Latent Dirichlet Allocation (LDA, Blei et al., 2003) model, we find a pretrained BERT model produces token embeddings that are more similar (in terms of inner product or cosine similarity) if they belong to the same topic, and more different if they belong to different topics (see e.g. Figure 3).

Inspired by these observations, we study LDA-generated data as a sandbox to understand—both through experiments on such synthetic data, and theoretical results—the process by which the embeddings and attention learn the topical structure. We find that the above observations from Wikipedia data are even more pronounced on synthetic LDA data. Moreover, we mathematically prove why such structure arises by analyzing a simplified two-stage training dynamics for a single-layer transformer trained under the masked language modeling objective. We also verify the two-stage nature of training dynamics obtains for a wide variety of optimizers and hyperparameter settings. Code is released at https://github.com/YuchenLi01/transformer_topic_model_LDA

OVERVIEW OF RESULTS

We focus on understanding the optimization dynamics of transformers in a simple sandbox: a single-layer transformer trained on (synthetic) data following a topic model distribution—and validate that our results robustly transfer to real data (Wikipedia WikimediaFoundation, 2023). We show that topic structure can be encoded both in the embedding layer, and in the attention mechanism of the network. Moreover, even if one of these components is not trained (i.e. handicapped), the other can “compensate” for it.

Theoretically, we characterize precisely how the topic structure is learned in the two extremal cases: when the attention mechanism is frozen to be uniform, and the only model parameters that are trained are the token embeddings; and when the token embeddings are frozen to be one-hot vectors, and the attention parameters (the key, query, and value matrices) are trained. We empirically verify our characterization on synthetic LDA-generated data, and also show that on real Wikipedia data, topic structure is learned both in the embeddings, and the attention mechanism.

In the first extremal case, we analyze the optima when we solely train the embedding layer. Precisely, we show that even when we freeze the attention scores to be uniform and all other elements of the transformer are set to identity, the model can still achieve near optimal loss by “encoding” the topic structure in the embedding weights:

Suppose the training data follows a topic model data distribution, and the transformer has trainable embedding layer, frozen (uniform) attention scores, and all other components set to identity. Then, the optimal embedding layer of a single layer transformer is such that the inner product of the embeddings of a pair of words is larger when the words belong to the same topic, and smaller when they belong to different topics.

Intuitively, this result states that words of the same topic, after training, have more similar embeddings than words of different topics. In this sense, the embedding layer captures the topic structure. We also empirically show (Section 6 and Figure 1) that this phenomenon is robust to differences in loss function and optimization method. See Section 4 for the formal theorem and Appendix B for the proof.

2 Topic structure is encoded in self-attention

In the second extreme, we study the behavior of the self-attention in a transformer trained on a topic modeling distribution, without the aid of trained token embeddings — i.e. when we use hard-coded, one-hot embeddings. The attention weight matrices WK{\bm{W}}^{K}, WQ{\bm{W}}^{Q}, and WV{\bm{W}}^{V} are initialized to near-zero matrices. To make the analysis feasible, we break down the training process into two separate stages, and characterize the optima in each stage. In the first stage, the attention is frozen to be uniform, and the matrix WV{\bm{W}}^{V} is trained. In the second stage, the matrix WV{\bm{W}}^{V} is frozen to the optimal value from the first stage, and the optimal attention weights is analyzed. Intuitively, such a two-stage approximation is reasonable, because in the initial stages of training, the gradients for the value matrix are much larger than those for the key and query matrices (see Section 8). While this is an approximation, this two-stage phenomenon can be observed empirically for a variety of hyperparameter settings (see Section 5.1 and in particular Figure 4). We also provide empirical evidence that the optima characterized in our analysis closely track the actual convergence points of models.

In brief, the self-attention function is Attn(Z)≔WVZA(Z)\text{Attn}(Z)\coloneqq{\bm{W}}^{V}{\bm{Z}}A({\bm{Z}}) in which A(Z)A({\bm{Z}}) denotes the attention weights, and WV{\bm{W}}^{V} is the value matrix weight. Intuitively, A(Z)ijA(Z)_{ij} is the importance of the i-th word for predicting the j-th word, and WV{\bm{W}}^{V} is aggregates the word embeddings in a sentence, weighted by the attention weights A(Z)A({\bm{Z}}). The formal definition of the model architecture is in Section 3.3.

We characterize the optimal WV{\bm{W}}^{V} in the initial stage of training: WV{\bm{W}}^{V} will learn a block-wise structure (see Figure 2), in which each block corresponds to a topic:

Suppose the training data follows a topic model data distribution, the token embeddings are frozen to be one-hot vectors, and attention scores are frozen to be uniform. Then, under mild L2L_{2} regularization, the optimal WV{\bm{W}}^{V} for the masked language modeling objective has block-wise structure, namely the (i,j)(i,j)-th entry of WV{\bm{W}}^{V} is on average larger when the tokens ii and jj belong to the same topic, and on average smaller when the tokens ii and jj belong to different topics.

For the formal theorem statement, see Section 5. The proof is deferred to Appendix D. We also empirically show (Section 6 and Figure 2) that this phenomenon is robust to differences in training loss and optimization method.

2.2 Optimal attention weights in Stage 2

For the second stage of the training dynamics, we assume WV{\bm{W}}^{V} is frozen to the optimal value in the first stage, and train the attention weights.

Suppose a single layer transformer is trained on a topic model data distribution, and WV{\bm{W}}^{V} is frozen to the block-wise first-stage optima. Then, the optimal attention weight for the masked language modeling objective is such that on average: a convex combination of same-word attention and same-topic-different-words attention should be relatively large, compared to different-topic attention.

For the formal assumption and theorem statements, see Section 5. The proof is deferred to Appendix D.

We empirically show (in Section 6) that even when the all the self-attention weight matrices are jointly trained (instead of trained with the two-stage process described), the behavior of attention weights still follows the relations that the above theorem describes.

3 Empirical results

We provide empirical evidence that the main conclusions in our theoretical findings remain robust even under settings that are more complex and realistic than our theoretical setup, and under variations of the training algorithm and loss. For example, we also test on synthetic data using a Latent Dirichlet Allocation (LDA) topic model (Blei et al., 2003) instead of our simplified topic modeling distribution; finally, we report results for a model pre-trained on the Wikipedia textual corpus, and discuss the connections with our conclusions derived in the synthetic setting. We describe detailed experimental setup and results in Section 6, as well as Appendix E.

PROBLEM SETUP

For our theoretical analysis, in order to have a well-defined notion of a “ground truth”, we will consider data distribution generated by a topic model consisting of TT topics {1,⋯ ,T}\{1,\cdots,T\} and TvTv words {1,⋯ ,Tv}\{1,\cdots,Tv\}. We will in fact, consider a special case of an LDA (Latent Dirichlet Allocation) model (Blei et al., 2003). Precisely, each document w{\bm{w}} is a sequence of words w1,⋯ ,wNw_{1},\cdots,w_{N}, and is generated by: Our theoretical results crucially depend on all topics being disjoint, i.e. they do not share common words. It is not crucial that the words in the same topic all have the same probabilities. Allowing these probabilities to be different would lead to results of similar flavor, but complicates the notation.

Randomly choose τ\tau distinct topics t1,⋯ ,tτt_{1},\cdots,t_{\tau} from [T][T].

Randomly choose a topic tt from {t1,⋯ ,tτ}\{t_{1},\cdots,t_{\tau}\}.

Randomly choose wnw_{n} from {(t−1)v+1,⋯ ,tv}\{{(t-1)v+1},\cdots,tv\}.

Note, under this data distribution, each word belongs to exactly one topic, and different topics do not share common words.

A word ii belongs to topic tt (denoted as i∈ti\in t) if i∈{(t−1)v+1,⋯ ,tv}i\in\{{(t-1)v+1},\cdots,tv\}. Correspondingly, topic(i)≔⌈iv⌉\texttt{topic}(i)\coloneqq\lceil\frac{i}{v}\rceil

Let Dw\mathcal{D}_{\bm{w}} denote the distribution of documents following the above generative process. Furthermore, for each document w{\bm{w}}, let X∈{0,1}(Tv+1)×N{\bm{X}}\in\{0,1\}^{(Tv+1)\times N} denote its one-hot encoding, in which Xij=1X_{ij}=1 if wj=iw_{j}=i, and 0 otherwise. Analogous to Dw\mathcal{D}_{\bm{w}}, let DX\mathcal{D}_{\bm{X}} denote the distribution of document one-hot encodings.

To simplify our theoretical analysis, we consider the infinitely-long-document setting, such that within each document, the empirical token distribution is equal to the groundtruth token distribution:

Each document w{\bm{w}} consists of exactly τ\tau topics {t1,⋯ ,tτ}\{t_{1},\cdots,t_{\tau}\}. Moreover, for each word i∈{1,⋯ ,Tv}i\in\{1,\cdots,Tv\} in the vocabulary, its empirical probability in the document

In our synthetic data experiments, we use a finite NN and generate data using an LDA model (Blei et al., 2003) which allows for slightly more variability—and demonstrates that our results are robust to changes in the setting. Detailed experimental setup is described in Section 6.

2 Training objective

Given data following the distribution defined in Section 3.1, we train a transformer network using the masked language modeling objective (Devlin et al., 2019). We first define the token [MASK]=0\texttt{[MASK]}=0 in addition to the words {1,⋯ ,Tv}\{1,\cdots,Tv\} of the topic model. Three constant probabilities pm,pc,pr∈(0,1)p_{m},p_{c},p_{r}\in(0,1) specify the masking scheme:

For the original document w=w1⋯wN{\bm{w}}=w_{1}\cdots w_{N}, first randomly choose a set of masked indices M(w)⊂[N]M({\bm{w}})\subset[N] such that ∀i∈[N]\forall i\in[N], with probability pmp_{m}, i∈M(w)i\in M({\bm{w}}).

Motivated by the empirical success of applying weight decay to training transformers, we also consider a regularized version of the above masked language modeling objective. For L2L_{2}-regularization When θ\theta is a vector, L2L_{2}-regularization penalizes ∥θ∥2\|\theta\|_{2}. When θ\theta is a matrix, the correct norm to regularize is ∥θ∥F\|\theta\|_{F}. with parameter λ>0\lambda>0:

Our experiments additionally study the cross entropy loss:

We give results for both types of loss functions because the cross-entropy loss, albeit practically more commonly used, is theoretically less convenient. Concretely, it involves the softmax operation which is invariant under addition by the same constant in each dimension (implying that the optimal logits are not necessarily unique); moreover, the optimal logits are often at infinity. By contrast, with squared loss, the set of optima is more easily characterized using some finite-valued closed form expressions.

Empirically, we will show (in Section 6) that the conclusions in our theoretical analyses hold for both the cross-entropy loss and the squared loss, as well as with variants of the training algorithm like SGD and Adam.

3 Transformer network architecture

To theoretically reason about the role played by the embedding layer and the self-attention layer, we consider a one-layer transformer model (Vaswani et al., 2017) with the simplification that the residual connection and normalization layers are removed. Precisely:

Appendix A includes additional remarks on the architecture.

In part of our theoretical analysis (in Section 5) and experiments (in Section 6), we freeze one-hot word embeddings, to study the mechanism that self-attention represents the topic structures without the aid of trained token embeddings. That is, set d=Tv+1d=Tv+1 and WE=I{\bm{W}}^{E}=I:

TOPIC STRUCTURE CAN BE ENCODED IN TOKEN EMBEDDINGS

The first result shows that, under the topic model data distribution, even if we freeze the self-attention to be uniform, the embedding layer can encode the topic structure. Precisely:

E00=−(1pm(1−pc−pr)−1)⋅u0{\bm{E}}_{00}=-\left(\frac{1}{p_{m}(1-p_{c}-p_{r})}-1\right)\cdot u_{0}

∀t∈[T],∑l∈tE0l=u0v\forall t\in[T],\sum_{l\in t}{\bm{E}}_{0l}=u_{0}v

The 0-th column of E{\bm{E}} satisfies ∀i∈{1,⋯ ,Tv}\forall i\in\{1,\cdots,Tv\}:

Ei0=−(1(1−pc−pr)pm−1)ui{\bm{E}}_{i0}=-\left(\frac{1}{(1-p_{c}-p_{r})p_{m}}-1\right)u_{i}

Eij{\bm{E}}_{ij} (∀i,j∈{1,⋯ ,Tv}\forall i,j\in\{1,\cdots,Tv\}) satisfy:

∑l∈topic(i)Eil=uiv+11−(1−pc)pm\sum_{l\in\texttt{topic}(i)}{\bm{E}}_{il}=u_{i}v+\frac{1}{1-(1-p_{c})p_{m}}

∀t∈[T]\forall t\in[T] such that topic(i)≠t\texttt{topic}(i)\neq t, ∑l∈tEil=uiv\sum_{l\in t}{\bm{E}}_{il}=u_{i}v

Point 3 is the important one among the list of conclusions. The way to read the theorem is that, among the entries of an optimal E{\bm{E}}: for ii and jj corresponding to the indices of tokens of the same topic, Eij{\bm{E}}_{ij} is (on average) larger, meaning that the embeddings of same-topic tokens are more similar; for ii and jj corresponding to different topics, Eij{\bm{E}}_{ij} is (on average) smaller, meaning that the embeddings of different-topic tokens are less similar. In particular, when the constants u0,⋯ ,uTvu_{0},\cdots,u_{Tv} are all zero, then the above larger-vs-smaller difference becomes a positive-vs-zero difference, which we roughly observe in practice.

Intuitively, the setting of the bias bpred{\bm{b}}^{\text{pred}} is used to “denoise” the masked sequence, i.e. to subtract the probability caused by filling in random words in the masking process (described in Section 3.2).

The proof of this theorem is deferred to Appendix B.

Proving comparable results under cross-entropy loss (equation 4) is more challenging considering Remark 1. However, we empirically show that, such blockwise pattern in E≔WE⊤WE{\bm{E}}\coloneqq{{\bm{W}}^{E}}^{\top}{\bm{W}}^{E} tends to exist in a trained model under both the squared loss and the cross-entropy loss, and regardless of whether we (i) train all layers or (ii) only train the embedding layer while freezing all other layers. Moreover, the loss achieved in case (ii) is only slightly worse than in case (i). Finally, we also show (Figure 3) that on real data, words that are unambiguous (e.g. “calculus”, “Mozart”) exhibit a similar pattern as Theorem 1 states: same-topic words have more similar embeddings, and therefore larger embedding dot products, than different-topic words. Quantitatively, if we only restrict ourselves to words that are unambigious (i.e. likely to be emitted only under few topics), a similar phenomenon can be observed (see Table 5).

TOPIC STRUCTURE CAN BE ENCODED IN SELF-ATTENTION

Whereas the previous section showed that the token embedding layer can in principle perform the heavy-lifting in learning the topic-modeling distribution, we further show that self-attention also can encode the topic structures, when we disallow training the embedding layer. That is, we freeze the token embeddings to be one-hot.

While inspecting the training dynamics of this one-layer transformer on the topic modeling data distribution, we observed a roughly two-stage process (illustrated by Figure 4): with certain initialization and learning rate settings, in Stage 1, the key matrix (WKW^{K}) and the query matrix (WQW^{Q}) stay close to 0, i.e. each position pays a near-uniform attention to all positions in the document, while the norm of the value matrix (WVW^{V}) increases significantly. In Stage 2, the norm of the the value matrix (WVW^{V}) already plateaus, and only after that, do the key and query matrices (WKW^{K} and WQW^{Q}) start to move.

Thus, while reasoning about the training process of transformers in our data distribution, we take motivation from the above empirical observation of such two-stage process, and consider a corresponding simplification: in Stage 1, the attention is frozen to be uniform, and only WV{\bm{W}}^{V} is trained; in Stage 2, WV{\bm{W}}^{V} is frozen, while WK{\bm{W}}^{K} and WQ{\bm{W}}^{Q} are trained. This simplification is a reasonable proxy for standard training, and we furthermore validate that our theoretical characterizations are robust to standard training, both using SGD and Adam. We provide more discussion on the two-stage optimization process in Section 8.

The Stage 1 of optimization process is convex (but not strongly convex) in WV{\bm{W}}^{V}, and we show that the set of minima consist of exactly the set of WV{\bm{W}}^{V} that exhibits a block-wise pattern:

∀j∈{0,⋯ ,Tv},W0jV∗=0\forall j\in\{0,\cdots,Tv\},{\bm{W}}^{V*}_{0j}=0

∀i∈{1,⋯ ,Tv},Wi0V∗=c2c3−c1Tvc22+Tv\forall i\in\{1,\cdots,Tv\},{\bm{W}}^{V*}_{i0}=\frac{c_{2}c_{3}-c_{1}Tv}{c_{2}^{2}+Tv}

WijV∗{\bm{W}}^{V*}_{ij} (∀i,j∈{1,⋯ ,Tv}\forall i,j\in\{1,\cdots,Tv\}):

∀l∉topic(i),  WilV∗=Wdiff-topicV∗≔−c1c2+c3c22+Tv\forall l\notin\texttt{topic}(i),\;{\bm{W}}^{V*}_{il}={\bm{W}}^{V*}_{\text{diff-topic}}\coloneqq-\frac{c_{1}c_{2}+c_{3}}{c_{2}^{2}+Tv}

∀l∈topic(i),WilV∗=Wsame-topicV∗≔Wdiff-topicV∗+c3v\forall l\in\texttt{topic}(i),{\bm{W}}^{V*}_{il}={\bm{W}}^{V*}_{\text{same-topic}}\coloneqq{\bm{W}}^{V*}_{\text{diff-topic}}+\frac{c_{3}}{v}

c1=pr(1−pc−pr)(1−(1−pc)pm)Tv∈(0,1)c_{1}=\frac{p_{r}}{(1-p_{c}-p_{r})\left(1-(1-p_{c})p_{m}\right)Tv}\in(0,1)

c2=1(1−pc−pr)pm−1∈(0,+∞)c_{2}=\frac{1}{(1-p_{c}-p_{r})p_{m}}-1\in(0,+\infty)

c3=11−(1−pc)pm∈(1,+∞)c_{3}=\frac{1}{1-(1-p_{c})p_{m}}\in(1,+\infty)

Empirically, the loss achieved by freezing WK=WQ=0{\bm{W}}^{K}={\bm{W}}^{Q}=0 and only training WV{\bm{W}}^{V} is only slightly greater than the loss achieved by training all of them jointly, see Appendix E.

Intuitively, this block-wise WV{\bm{W}}^{V} shows that, while inferring about the words at the masked positions: the model looks at unmasked positions in the document, each unmasked word only contributes to predicting words of the same topic, each unmasked word does not contribute to predicting words of different topics, and the model implicitly aggregates the topic distribution among the unmasked words, to infer the token distribution in the original document prior to masking.

The proof of this Theorem 2 is deferred to Appendix C. Proving a comparable result under the cross-entropy loss equation 4 is more challenging due to the same reasons outlined in Remark 1. However, empirically such block-wise WV{\bm{W}}^{V} shows up for both the cross-entropy loss and the squared loss, as we show in Section 6.

3 Optimal attention weights

In our analysis on the stage 2 optimization process, we freeze the WV{\bm{W}}^{V} to be some representative optima from stage 1 (Theorem 2), and characterize the optimal attention weights by comparing the following three types of attention weights: among the same words at different positions, among different words of the same topic, and among words of different topics.

We mainly consider the type of optimal WV{\bm{W}}^{V} characterized in Theorem 2: WV{\bm{W}}^{V} with uniform blocks (see Figure 2). Empirically, the model often approximately converges to these type of pattern (Section 6).

To formally reason about the behavior of average attention weights, we consider a simplified setting:

in which c2=αc3c_{2}=\alpha c_{3} and c1=βc3c_{1}=\beta c_{3}.

We will characterize the setting of α\alpha and β\beta that minimizes the loss, under the following assumptions:

T→∞T\to\infty, i.e. the total number of topics grows to infinity.

(Sparse documents): τ→∞,τ=o(T)\tau\to\infty,\tau=o(T), i.e. the number of topics in each document also grows to infinity, but much smaller than the total number of topics. (This is a common parameter regime: we typically think of each document as a sparse combination of topics.)

(No sparsely supported topics): v>(11−(1−pc)pm+1)2+1v>(\frac{1}{1-(1-p_{c})p_{m}}+1)^{2}+1 (vv is the number of tokens in each topic. v≥10v\geq 10 suffices under Assumption 4. This is also a common regime, where we assume no topic consists only of a small number of words.)

In the training objective (Section 3.2), we consider the case pm<12,  pc=pr∈(0,12)p_{m}<\frac{1}{2},\;p_{c}=p_{r}\in(0,\frac{1}{2}). This setting is consistent with the masking scheme proposed in Devlin et al. (2019).

Suppose the data distribution follows the topic modeling assumption in Section 3.1 and Assumption 1. Suppose we train a single layer transformer given by equation 7 with bpred=0{\bm{b}}^{\text{pred}}=0 and WV{\bm{W}}^{V} frozen to the optima in Theorem 2, under masked language modeling objective (equation 1) with the squared loss (equation 3), under Assumption 2, Assumption 3, and Assumption 4. Then, the optimal (α,β)(\alpha,\beta) satisfy

in which λ1≔(1−(1−pc)pm+pmpr)(1+(1−pc)pm)2(1−(1−pc)pm)\lambda_{1}\coloneqq\frac{(1-(1-p_{c})p_{m}+p_{m}p_{r})(1+(1-p_{c})p_{m})}{2(1-(1-p_{c})p_{m})} and λ2≔100(1−(1−pc)pmpmpr+1)\lambda_{2}\coloneqq 100(\frac{1-(1-p_{c})p_{m}}{p_{m}p_{r}}+1).

In particular, Theorem 3 implies that if we choose τ,T\tau,T such that the lower bound exceeds 1, we expect the attention between same-topic words to be on average larger than that between different-topic words.

Note that when WV{\bm{W}}^{V} is block-diagonal with uniform blocks, it is impossible to meaningfully bound α\alpha or β\beta individually; instead, only their weighted average (v−1vα+1vβ\frac{v-1}{v}\alpha+\frac{1}{v}\beta) matters. In other words, different (α,β)(\alpha,\beta) will incur the same loss, as long as the above weighted average remains the same. Intuitively, this is because such block-diagonal WV{\bm{W}}^{V} with uniform blocks sums up the attention on all words in each topic, and make predictions solely based on the sums. The proof of Theorem 3 is deferred to Appendix D.3.

When there is no L2L_{2}-regularization, the first-stage optima of WV{\bm{W}}^{V} is not unique. We include additional analysis for representative cases of WV{\bm{W}}^{V} in Appendix D.4.

When T,τT,\tau are finite, the loss expression turns out to be too complicated to characterize in closed form (because all the o(1)o(1) terms need to be expanded). So we instead numerically compute the loss landscape as a function of α\alpha and β\beta. See Appendix D.5.

EXPERIMENTS

We analyze properties of the training dynamics via extensive experimental analysis. We will describe both the setup for synthetic (LDA-generated) data, and for Wikipedia data.

In our experiments, we generate data following Section 3.1 with T=10,v=10T=10,v=10, NN uniformly randomly chosen from $,exceptthatStep1ischangedtosamplingthetopicdistributionaccordingtotheDirichletdistribution(consistentwithLDA,Bleietal.,2003)with, except that Step 1 is changed to sampling the topic distribution according to the Dirichlet distribution (consistent with LDA, Blei et al., 2003) with\alpha=0.1.Mostsentencescontain2to4topics.OurtrainingobjectivefollowsSection3.2with. Most sentences contain 2 to 4 topics. Our training objective follows Section 3.2 withp_{m}=0.15,p_{c}=0.1,p_{r}=0.1followingDevlinetal.(2019).WeusethemodelarchitecturefollowingSection3.3butaddbackthebiastermsfollowing Devlin et al. (2019). We use the model architecture following Section 3.3 but add back the bias terms{\bm{b}}^{K},{\bm{b}}^{Q},{\bm{b}}^{V}$, following standard implementation in Wolf et al. (2020).

In Figure 1, we show that for a model in which all components are trained, the learned embedding weight WE{\bm{W}}^{E} is such that WE⊤WE{{\bm{W}}^{E}}^{\top}{\bm{W}}^{E} displays a block-wise pattern. In particular, a diagonal pattern is a special case. These results show that our theory in Section 4 characterizes the optima of embedding layer which can be found by using either cross-entropy or squared losses, either SGD or Adam optimizers, and even when the other layers in the model are trained instead of frozen.

We show that when the word embeddings are frozen to one-hot and the attention weights are uniform (by setting WK=0,WQ=0{\bm{W}}^{K}=0,{\bm{W}}^{Q}=0), the trained WV{\bm{W}}^{V} has a block-wise pattern, corresponding to the topical structure (see Figure 2).

We show (in Figure 10 in Appendix E.1) that even when the attention weights WK,WQ{\bm{W}}^{K},{\bm{W}}^{Q} are jointly trained with WV{\bm{W}}^{V}, the model would still approximately converge to the type of block-wise WV{\bm{W}}^{V} described in our analyses in Section 5.2.

We show that, our conclusion in Theorem 3 holds not just when WV{\bm{W}}^{V} is frozen to a block-wise pattern, but also when it is trained and naturally converges to such pattern. And we show (in Table 3 in Appendix E.2) that on average, each word pays more attention to words of the same topic than to words of different topics.

2 Results on natural language data

For a set of pre-trained transformer-based models (and their corresponding tokenizers) downloaded from Huggingface (Wolf et al., 2020), we compare the embedding similarity and attention weights between same-topic tokens and different-topic tokens. The topics are determined by fitting an LDA model with 100 topics on a sample of Wikipedia corpus (WikimediaFoundation, 2023) tokenized by the above tokenizers. We filter stop words. For each topic, we only keep a fraction of tokens that LDA assigns the highest likelihood in this topic. Consistent with our theoretical setting, we restrict to keeping only one topic for each word. In Table 1, we provide the results after such pre-processing. We provide additional details about the experimental setup and additional results (including when the last restriction of “one topic per word” is removed) in Appendix E.3.

RELATED WORKS

One line of prior works explain the success of transformers by empirically showing that the components (e.g. attention heads) of a trained model (e.g. BERT Devlin et al., 2019), contain abundant information for solving a wide range of “probing” tasks, across syntax and semantics (Hewitt & Manning, 2019; Clark et al., 2019; Tenney et al., 2019; Hewitt & Liang, 2019; Kovaleva et al., 2019; Belinkov, 2022), or through other approaches involving the attention weights (Vig & Belinkov, 2019; Htut et al., 2019; Sun & Marasović, 2021). Our result also formalizes some relevant intuitions given in Elhage et al. (2021), such as embedding layer capturing some bigram statistics. In topic modeling distribution, such “bigram statistics” translates to co-occurrence in a document.

Recent works start to combine theoretical constructions and controlled experiments to justify the expressive power of transformers through the lens of Turing completeness (Bhattamishra et al., 2020b), function approximation (Yun et al., 2020), representing formal languages (Bhattamishra et al., 2020a; Ebrahimi et al., 2020; Yao et al., 2021; Liu et al., 2023), learning abstract algebraic operations (Zhang et al., 2022a), statistical sample complexity (Wei et al., 2021; Edelman et al., 2022), and learning optimal latent representation (Zhang et al., 2023). Methodologically, we join a long line of works that characterize the capacity of neural network models by assessing their abilities in learning some simple models of the data (Siegelmann & Sontag, 1992; Gers & Schmidhuber, 2001; Weiss et al., 2018; Suzgun et al., 2019; Merrill, 2019; Hewitt et al., 2020; Li & Risteski, 2021; Yao et al., 2021; Zhang et al., 2022a; Liu et al., 2023). Our work extends this line of works, and in particular, our results indicate that there may be multiple reasonable representational optima, which calls for formally analyzing the training dynamics to gain deeper understanding of what the model actually learns from such data distributions.

On the optimization side, Nguyen & Salazar (2019); Xiong et al. (2020); Liu et al. (2020); Zhang et al. (2020); Li & Gong (2021) propose algorithmic improvements (often with theoretical motivations) to help stabilize the training process of transformers. Towards explaining the training process of attention-based neural networks, Sun & Lu (2020) analyzes the trends of two quantities that are relevant to model performance and interpretability in text classification setting.

Also relevant to our work, Snell et al. (2021) consider cross-attention in LSTM Seq2Seq models trained on machine-translation settingsSpecifically, they consider a data model related to the IBM machine translation model.. By contrast, we focus on self-attention in transformers, and we consider a data distribution inspired by topic models. Notably, they also propose an intuitive simplifying assumption of a two-stage learning process of the attention heads similar to ours (but without theoretical or empirical validation). Our work uses a similar assumption We independently proposed the two-stage training of attention heads, and later discovered (Snell et al., 2021) used a similar assumption. Comparison with (Snell et al., 2021) was added during an update of our paper. Moreover, while Snell et al. (2021) is the earliest paper we are aware of that explicitly assumes a two-stage training process specifically for attention heads, we note that similar approaches (more generally, alternating optimization) commonly appear in the optimization literature in a broad variety of settings. (Section 5.1). In our work, we validate our version of the two-stage assumption by providing a particular way to initialize the attention weight matrices, along with theoretical intuitions (Section 8) and empirical validation on synthetic data (Figure 4) as well as real data (Figure 5), showing that this two-stage process can be a reasonable approximation to the early steps of the real training dynamics of attention-based models under the settings that we analyze.

Recent work by Jelassi et al. (2022) theoretically shows how transformers learn the spatial structure of image-type datasets through gradient-descent-based optimization algorithms. In particular, their attention weights depend on the positional encodings only. Different from their work, our result (motivated by studying the semantics in language) focuses on topic modeling distribution that actually ignores the position information, so the attention weights only depend on the “bag of words” (i.e. the contents). In that sense, Jelassi et al. (2022) and our work complement each other, since real-world data distribution usually involves a combination of position-dependent and position-independent factors. An interesting future work would be studying how these factors interact during the training process.

Regarding the type of data distribution that we consider, we join a series of works that theoretically reason about the ability of learning under topic-modeling-based distributions (Sontag & Roy, 2011; Awasthi & Risteski, 2015; Arora et al., 2016; Tosh et al., 2021; Luo et al., 2022). In particular, Luo et al. (2022) shows that if a model can achieve low loss on contrastive or mask-prediction objectives, then it can recover topic posterior. However, these prior works do not theoretically analyze the optimization process of the transformer architecture. In fact, model architecture can indeed critically influence the resulting model obtained by masked-prediction-type tasks (see Liu et al. (2022) who highlight the subtlety of the interaction between the particular form of the task and the model specification). Hence, our analysis extends beyond the scope of these prior works by incorporating the theoretical analysis on the optimization process of transformers trained on topic modeling data distribution. Empirically, Sia et al. (2020); Thompson & Mimno (2020); Meng et al. (2022); Zhang et al. (2022b); Talebpour et al. (2023) analyze topic discovery via clustering the contextualized representations produced by pretrained language models. Different from these works, our theory and experiments on token embeddings focus on the convergence of embedding layer parameters.

DISCUSSION

This two-stage optimization process (Section 5.1 and Figure 4) can be thought of as one iteration of the alternating optimization procedure. That is, we first train WV{\bm{W}}^{V} while freezing (WK,WQ)({\bm{W}}^{K},{\bm{W}}^{Q}), and then freeze WV{\bm{W}}^{V} while training (WK,WQ)({\bm{W}}^{K},{\bm{W}}^{Q}), and repeat this process.

In practice, WK,WQ,WV{\bm{W}}^{K},{\bm{W}}^{Q},{\bm{W}}^{V} in transformers are typically trained jointly instead of alternatingly. However, our empirical results show that, the conclusions drawn from the two-stage optimization analysis carry over even when they are trained jointly. Moreover, we don’t find any qualitative aspects of normal training that are not captured by this two-stage approximation.

Therefore, in the initial steps (i.e. Stage 1), WV{\bm{W}}^{V} intuitively grows much faster than WK{\bm{W}}^{K}. For the same reason (note the symmetry between WK{\bm{W}}^{K} and WQ{\bm{W}}^{Q}, see equation 5), WV{\bm{W}}^{V} intuitively grows much faster than WQ{\bm{W}}^{Q}, too.

In Stage 2, it is less intuitively clear why ∥WV∥F\|{\bm{W}}^{V}\|_{F} tends to plateau. Note that empirically, even when ∥WV∥F\|{\bm{W}}^{V}\|_{F} plateaus, the WV{\bm{W}}^{V} matrix itself still fluctuates with non-vanishing step-by-step changes. (That is, in each step, WV{\bm{W}}^{V} “locally rotates” around the origin with an approximately constant norm.) Hence we refer to our Stage 2 analysis (which freezes WV{\bm{W}}^{V} itself) as a simplification. However, the final empirical convergence point of WV{\bm{W}}^{V} matches our theoretical analysis.

We show in Figure 5 that an approximate version of this multi-stage phenomenon can be observed on multi-layer transformers trained on Wikipedia as well.

Finally, this two-stage phenomenon is sensitive to hyperparameters like initialization and learning rate. In Figure 4, the The training process is not usually visibly two-stage using the common default hyperparameters. We leave it as an interesting future work to theoretically analyze the training dynamics when the two-stage phenomenon is not present.

2 Do topic-wise behaviors perfectly correlate with co-occurrence counts?

Additionally, we note that fitting a topic model is closely related to word co-occurrence statistics, which raises the following question: should those empirical phenomenon (i.e. higher same-topic attention and more similar same-topic embeddings, shown in Table 5) be more fundamentally attributed to larger co-occurrence counts?

In the following, we also compare them with some preliminary empirical results on the behavior of embedding and attention, from both topic modeling and co-occurrence perspectives. Specifically, we compare the average attention weights and average embedding dot products, between same-topic word pairs and the NN pairs of words that co-occur the most frequently in a sample of the Wikipedia corpus. The cutoff NN is determined so that the number of ”top co-occurring word pairs” is the same as the number of word pairs in each topic (controlled by the ambiguity threshold). The results are summarized in Table 2.

Based on those results, we conjecture that the topic-wise behavior of token embeddings and attention weights cannot be fully explained by simple co-occurrence counts.

Reasoning about their connections more formally would require analyzing some data distributions that better decouple these factors. We think that would be an interesting direction of future work.

CONCLUSION

We initiated the study of understanding training dynamics of transformers in the presence of semantic structure captured by a topic model. Interesting directions of future work includes extending the analysis to data distributions that captures “syntactic” structure, e.g. through simple sandboxes like PCFGs. When both the model and the data distributions are complex, it remains a daunting challenge to “disentangle” how the many different aspects of the data (e.g. semantic and syntactic elements) are learned through the different parts of model architecture (e.g. attention, positional encodings, and embeddings).

We thank Bingbin Liu, Yusha Liu, and Tanya Marwah for proofreading and providing constructive comments, Yewen Fan for helpful suggestions on empirically obtaining the two-stage optimization process, and Emmy Liu and Graham Neubig for insightful discussions on the connections with empirical observations.

Andrej Risteski and Yuchen Li acknowledge support by NSF awards IIS-2211907 and CCF-2238523. Andrej Risteski also acknowledges support by Amazon Research Award “Causal + Deep Out-of-Distribution Learning”.

References

Appendix A ADDITIONAL INFORMATION ON THE SETUP

The positional encoding at the input is also removed, because the position information of a word in a document is irrelevant to the topic model defined in Section 3.1.

Under our setting, we first prove the following useful Lemma 1. Intuitively, it states that, when freezing uniform attention, the output of self-attention weights essentially counts the unmasked tokens in the document (as a result of the masking process described in Section 3.2). Given those counts, the best way to predict a token at the masked positions in the original document (i.e. prior to the masking process) is to:

First, aggregate the counts of the unmasked words within each topic, to infer the topic distribution in the observed document. In this, we further have the restriction that:

Each unmasked word only contributes to predicting words of the same topic

Each unmasked word does not contribute to predicting words of different topics

Never predict the mask token ([MASK]), because the original document does not contain any [MASK]

Second, we “denoise” the topic distribution, i.e. we subtract the probability caused by filling in random words in the masking process (described in Section 3.2).

W00=−(1pm(1−pc−pr)−1)⋅u0{\bm{W}}_{00}=-\left(\frac{1}{p_{m}(1-p_{c}-p_{r})}-1\right)\cdot u_{0}

∀t∈[T],∑l∈tW0l=u0v\forall t\in[T],\sum_{l\in t}{\bm{W}}_{0l}=u_{0}v

∀i∈{1,⋯ ,Tv},Wi0=−pr(1−pc−pr)(1−(1−pc)pm)Tv−(1(1−pc−pr)pm−1)ui\forall i\in\{1,\cdots,Tv\},{\bm{W}}_{i0}=-\frac{p_{r}}{(1-p_{c}-p_{r})\left(1-(1-p_{c})p_{m}\right)Tv}-\left(\frac{1}{(1-p_{c}-p_{r})p_{m}}-1\right)u_{i}

Wij{\bm{W}}_{ij} (∀i,j∈{1,⋯ ,Tv}\forall i,j\in\{1,\cdots,Tv\}):

∑l∈topic(i)Wil=11−(1−pc)pm+uiv\sum_{l\in\texttt{topic}(i)}{\bm{W}}_{il}=\frac{1}{1-(1-p_{c})p_{m}}+u_{i}v

∀t∈[T]\forall t\in[T] such that topic(i)≠t\texttt{topic}(i)\neq t, ∑l∈tWil=uiv\sum_{l\in t}{\bm{W}}_{il}=u_{i}v

In fact, this optimization objective becomes strongly convex with an L2L_{2} regularization for some λ>0\lambda>0.

and the last step follows since ∀l∈topic(i),Pw(l)=Pw(i)\forall l\in\texttt{topic}(i),P_{{\bm{w}}}(l)=P_{{\bm{w}}}(i) under our setting in Section 3.1.

and so L(W)L({\bm{W}}) is minimized when ∀X\forall{\bm{X}},

which requires ∀i∈{0,⋯ ,Tv+1}\forall i\in\{0,\cdots,Tv+1\},

From equation A.9 and equation A.10 we get:

Note that under the topic modeling distribution in Section 3.1, for any topic t∈[T]t\in[T],

Hence we simplify equation A.11 by considering the proportions of the “representative” tokens for each topic:

We obtain: for all sets of {Pw(i):i∈[Tv]}\{P_{{\bm{w}}}(i):i\in[Tv]\} satisfying our distribution in Section 3.1

Specifically, fix Pw(i)=12vP_{{\bm{w}}}(i)=\frac{1}{2v} and consider the following settings of {Pw(j):j∉topic(i)}\{P_{{\bm{w}}}(j):j\notin\texttt{topic}(i)\}:

Pw(j)=12vP_{{\bm{w}}}(j)=\frac{1}{2v} if topic(j)=t1\texttt{topic}(j)=t_{1} and 0 otherwise. Then equation A.13 becomes

Pw(j)=12vP_{{\bm{w}}}(j)=\frac{1}{2v} if topic(j)=t2\texttt{topic}(j)=t_{2} and 0 otherwise. Then equation A.13 becomes

Clearly the above two equations cannot both hold, because ∑l∈t1Wil>∑l∈t2Wil\sum_{l\in t_{1}}{\bm{W}}_{il}>\sum_{l\in t_{2}}{\bm{W}}_{il}.

Hence we proved by contradiction that ∀t1,t2≠topic(i),∑l∈t1Wil=∑l∈t2Wil\forall t_{1},t_{2}\neq\texttt{topic}(i),\sum_{l\in t_{1}}{\bm{W}}_{il}=\sum_{l\in t_{2}}{\bm{W}}_{il}. Likewise, when i=0i=0, ∀t1,t2in[T],∑l∈t1W0l=∑l∈t2W0l\forall t_{1},t_{2}in[T],\sum_{l\in t_{1}}{\bm{W}}_{0l}=\sum_{l\in t_{2}}{\bm{W}}_{0l}.

Since this has to hold for all Pw(i)∈[0,1v]P_{{\bm{w}}}(i)\in[0,\frac{1}{v}], the coefficients must match, i.e.

Appendix B PROOF OF THEOREM 1: OPTIMAL TOKEN EMBEDDING

Consider training a transformer given by equation 6 with WK=0,WQ=0,WV=I{\bm{W}}^{K}=0,{\bm{W}}^{Q}=0,{\bm{W}}^{V}=I and ∀i∈{1,⋯ ,Tv},bipred=−pmpr(1−(1−pc)pm)Tv\forall i\in\{1,\cdots,Tv\},{\bm{b}}^{\text{pred}}_{i}=-\frac{p_{m}p_{r}}{\left(1-(1-p_{c})p_{m}\right)Tv} on data coming from the topic model described in Section 3, with the masked language modeling objective (equation 1) with squared loss (equation 3).

E00=−(1pm(1−pc−pr)−1)⋅u0{\bm{E}}_{00}=-\left(\frac{1}{p_{m}(1-p_{c}-p_{r})}-1\right)\cdot u_{0}

∀t∈[T],∑l∈tE0l=u0v\forall t\in[T],\sum_{l\in t}{\bm{E}}_{0l}=u_{0}v

∀i∈{1,⋯ ,Tv},Ei0=−(1(1−pc−pr)pm−1)ui\forall i\in\{1,\cdots,Tv\},{\bm{E}}_{i0}=-\left(\frac{1}{(1-p_{c}-p_{r})p_{m}}-1\right)u_{i}

Eij{\bm{E}}_{ij} (∀i,j∈{1,⋯ ,Tv}\forall i,j\in\{1,\cdots,Tv\}):

∑l∈topic(i)Eil=11−(1−pc)pm+uiv\sum_{l\in\texttt{topic}(i)}{\bm{E}}_{il}=\frac{1}{1-(1-p_{c})p_{m}}+u_{i}v

∀t∈[T]\forall t\in[T] such that topic(i)≠t\texttt{topic}(i)\neq t, ∑l∈tEil=uiv\sum_{l\in t}{\bm{E}}_{il}=u_{i}v

and the last step is because by equation D.17,

Let E′∗{\bm{E}}^{\prime*} denote any matrix in

E00′∗=−(1pm(1−pc−pr)−1)⋅u0{\bm{E}}^{\prime*}_{00}=-\left(\frac{1}{p_{m}(1-p_{c}-p_{r})}-1\right)\cdot u_{0}

∀t∈[T],∑l∈tE0l′∗=u0v\forall t\in[T],\sum_{l\in t}{\bm{E}}^{\prime*}_{0l}=u_{0}v

∀i∈{1,⋯ ,Tv},Ei0′∗=−pr(1−pc−pr)(1−(1−pc)pm)Tv−(1(1−pc−pr)pm−1)ui\forall i\in\{1,\cdots,Tv\},{\bm{E}}^{\prime*}_{i0}=-\frac{p_{r}}{(1-p_{c}-p_{r})\left(1-(1-p_{c})p_{m}\right)Tv}-\left(\frac{1}{(1-p_{c}-p_{r})p_{m}}-1\right)u_{i}

Eij′∗{\bm{E}}^{\prime*}_{ij} (∀i,j∈{1,⋯ ,Tv}\forall i,j\in\{1,\cdots,Tv\}):

∑l∈topic(i)Eil′∗=11−(1−pc)pm+uiv\sum_{l\in\texttt{topic}(i)}{\bm{E}}^{\prime*}_{il}=\frac{1}{1-(1-p_{c})p_{m}}+u_{i}v

∀t∈[T]\forall t\in[T] such that topic(i)≠t\texttt{topic}(i)\neq t, ∑l∈tEil′∗=uiv\sum_{l\in t}{\bm{E}}^{\prime*}_{il}=u_{i}v

Therefore, by equation B.16, let E∗{\bm{E}}^{*} denote any matrix in

E00∗=−(1pm(1−pc−pr)−1)⋅u0{\bm{E}}^{*}_{00}=-\left(\frac{1}{p_{m}(1-p_{c}-p_{r})}-1\right)\cdot u_{0}

∀t∈[T],∑l∈tE0l∗=u0v\forall t\in[T],\sum_{l\in t}{\bm{E}}^{*}_{0l}=u_{0}v

∀i∈{1,⋯ ,Tv},Ei0∗=−(1(1−pc−pr)pm−1)ui\forall i\in\{1,\cdots,Tv\},{\bm{E}}^{*}_{i0}=-\left(\frac{1}{(1-p_{c}-p_{r})p_{m}}-1\right)u_{i}

Eij∗{\bm{E}}^{*}_{ij} (∀i,j∈{1,⋯ ,Tv}\forall i,j\in\{1,\cdots,Tv\}):

∑l∈topic(i)Eil∗=11−(1−pc)pm+uiv\sum_{l\in\texttt{topic}(i)}{\bm{E}}^{*}_{il}=\frac{1}{1-(1-p_{c})p_{m}}+u_{i}v

∀t∈[T]\forall t\in[T] such that topic(i)≠t\texttt{topic}(i)\neq t, ∑l∈tEil∗=uiv\sum_{l\in t}{\bm{E}}^{*}_{il}=u_{i}v

W00V=−(1pm(1−pc−pr)−1)⋅u0{\bm{W}}^{V}_{00}=-\left(\frac{1}{p_{m}(1-p_{c}-p_{r})}-1\right)\cdot u_{0}

∀t∈[T],∑l∈tW0lV=u0v\forall t\in[T],\sum_{l\in t}{\bm{W}}^{V}_{0l}=u_{0}v

∀i∈{1,⋯ ,Tv},Wi0V=−pr(1−pc−pr)(1−(1−pc)pm)Tv−(1(1−pc−pr)pm−1)ui\forall i\in\{1,\cdots,Tv\},{\bm{W}}^{V}_{i0}=-\frac{p_{r}}{(1-p_{c}-p_{r})\left(1-(1-p_{c})p_{m}\right)Tv}-\left(\frac{1}{(1-p_{c}-p_{r})p_{m}}-1\right)u_{i}

WijV{\bm{W}}^{V}_{ij} (∀i,j∈{1,⋯ ,Tv}\forall i,j\in\{1,\cdots,Tv\}):

∑l∈topic(i)WilV=11−(1−pc)pm+uiv\sum_{l\in\texttt{topic}(i)}{\bm{W}}^{V}_{il}=\frac{1}{1-(1-p_{c})p_{m}}+u_{i}v

∀t∈[T]\forall t\in[T] such that topic(i)≠t\texttt{topic}(i)\neq t, ∑l∈tWilV=uiv\sum_{l\in t}{\bm{W}}^{V}_{il}=u_{i}v

Note that this is exactly the statement of Lemma 1 (proved in Appendix A.1) in the case of W≔WV{\bm{W}}\coloneqq{\bm{W}}^{V}. ∎

∀j∈{0,⋯ ,Tv},W0jV∗=0\forall j\in\{0,\cdots,Tv\},{\bm{W}}^{V*}_{0j}=0

∀i∈{1,⋯ ,Tv},Wi0V∗=c2c3−c1Tvc22+Tv\forall i\in\{1,\cdots,Tv\},{\bm{W}}^{V*}_{i0}=\frac{c_{2}c_{3}-c_{1}Tv}{c_{2}^{2}+Tv}

WijV∗{\bm{W}}^{V*}_{ij} (∀i,j∈{1,⋯ ,Tv}\forall i,j\in\{1,\cdots,Tv\}):

∀l∉topic(i),  WilV∗=Wdiff-topicV∗≔−c1c2+c3c22+Tv\forall l\notin\texttt{topic}(i),\;{\bm{W}}^{V*}_{il}={\bm{W}}^{V*}_{\text{diff-topic}}\coloneqq-\frac{c_{1}c_{2}+c_{3}}{c_{2}^{2}+Tv}

∀l∈topic(i),  WilV∗=Wsame-topicV∗≔Wdiff-topicV∗+c3v\forall l\in\texttt{topic}(i),\;{\bm{W}}^{V*}_{il}={\bm{W}}^{V*}_{\text{same-topic}}\coloneqq{\bm{W}}^{V*}_{\text{diff-topic}}+\frac{c_{3}}{v}

c1=pr(1−pc−pr)(1−(1−pc)pm)Tv∈(0,1)c_{1}=\frac{p_{r}}{(1-p_{c}-p_{r})\left(1-(1-p_{c})p_{m}\right)Tv}\in(0,1)

c2=1(1−pc−pr)pm−1∈(0,+∞)c_{2}=\frac{1}{(1-p_{c}-p_{r})p_{m}}-1\in(0,+\infty)

c3=11−(1−pc)pm∈(1,+∞)c_{3}=\frac{1}{1-(1-p_{c})p_{m}}\in(1,+\infty)

Step 1: the optima converges to one outlined in Lemma 1

In comparison, ∀W∈S\forall{\bm{W}}\in S, by Lemma 1, since WV∗∉S{\bm{W}}^{V*}\notin S,

Moreover, note that since ∥W∥F\|{\bm{W}}\|_{F} is finite,

Combining the above two observations gives

Therefore, we have proved by contradiction that

Step 2: solve for the coefficients that minimize the L2L_{2} penalty

in which the last step is because ∀WV∈S\forall{\bm{W}}^{V}\in S, L(WV)=min⁡L(WV)L({\bm{W}}^{V})=\min L({\bm{W}}^{V}), which is a constant independent of WV{\bm{W}}^{V}.

Appendix D ADDITIONAL RESULTS ON ATTENTION WEIGHTS

In this section, we will calculate a few expressions for the masking probabilities, which will be useful for the proofs later on. We will also introduce a few constants for brevity of notation.

A straightforward calculation shows that the probabilities after the masking process satisfy:

For convenience, we will introduce the notation

Another straightforward calculation can be used to express the relationship between the constant c3c_{3} in Assumption 2 and the α,β\alpha,\beta. Namely, we have:

The constant c3c_{3} in Assumption 2 satisfies:

Again, for notational convenience, we will introduce z1,z2z_{1},z_{2}, s.t.

D.2 Implication of topic-wise attention assumption on model output

Using the calculation and notations of z1z_{1} and z2z_{2} in Appendix D.1:

Suppose the data distribution follows the topic modeling assumption in Section 3.1 and Assumption 1. Suppose we train a single layer transformer given by equation 7 with bpred=0{\bm{b}}^{\text{pred}}=0 and WV{\bm{W}}^{V} frozen to the optima in Theorem 2, under masked language modeling objective (equation 1) with the squared loss (equation 3), under Assumption 2, Assumption 3, and Assumption 4. Then, the optimal (α,β)(\alpha,\beta) satisfy

in which the constants λ1≔(1−(1−pc)pm+pmpr)(1+(1−pc)pm)2(1−(1−pc)pm)\lambda_{1}\coloneqq\frac{(1-(1-p_{c})p_{m}+p_{m}p_{r})(1+(1-p_{c})p_{m})}{2(1-(1-p_{c})p_{m})} and λ2≔100(1−(1−pc)pmpmpr+1)\lambda_{2}\coloneqq 100(\frac{1-(1-p_{c})p_{m}}{p_{m}p_{r}}+1).

Define γ≔v−1vα+1vβ\gamma\coloneqq\frac{v-1}{v}\alpha+\frac{1}{v}\beta.

Recall the architecture under consideration, i.e.

For a document w{\bm{w}} which contains topics t1,⋯ ,tτ∈[T]t_{1},\cdots,t_{\tau}\in[T], there are the following cases:

Plugging in the asymptotics from Assumption 3, the above becomes

Plugging in the asymptotics from Assumption 3, the above terms vanish.

Plugging in the asymptotics from Assumption 3, the above becomes

Adding equation D.23 and equation D.26, we can see in the asymptotic regime of interest, we have:

in which the constant c4c_{4} is defined as

c4≔11−(1−pc)pm∈(1,2)c_{4}\coloneqq\frac{1}{1-(1-p_{c})p_{m}}\in(1,2)

Plugging in the definition of p1,p2p_{1},p_{2} in equation D.18, equation D.19

We will again consider several possible cases for γ\gamma in equation D.27.

Case 1: When γ≤(1+c4pmpr)(2−c4)2c4(τ−1)\gamma\leq\frac{(1+c_{4}p_{m}p_{r})(2-c_{4})}{2c_{4}}(\tau-1).

then focusing on this term in the loss equation D.27:

Case 2: When γ≥1001c4+pmprpmprT\gamma\geq 100\frac{\frac{1}{c_{4}}+p_{m}p_{r}}{p_{m}p_{r}}T.

and therefore plugging into equation D.27:

Case 3: When (1+c4pmpr)(2−c4)2c4(τ−1)<γ<1001c4+pmprpmprT\frac{(1+c_{4}p_{m}p_{r})(2-c_{4})}{2c_{4}}(\tau-1)<\gamma<100\frac{\frac{1}{c_{4}}+p_{m}p_{r}}{p_{m}p_{r}}T.

then similar to Case 2, since τ=o(T)\tau=o(T) by Assumption 3:

and therefore plugging into equation D.27:

Note that L(γ)L(\gamma) in Case 3 is strictly smaller than L(γ)L(\gamma) in Case 1 and Case 2, because:

Comparing equation D.28 and equation D.30: (1−c4v(1+1+c4pmprc5))2>1−c4(2−c4)v(1-\frac{c_{4}}{v(1+\frac{1+c_{4}p_{m}p_{r}}{c_{5}})})^{2}>1-\frac{c_{4}(2-c_{4})}{v} because c5∈(0,(1+c4pmpr)(2−c4)c4)c_{5}\in(0,\frac{(1+c_{4}p_{m}p_{r})(2-c_{4})}{c_{4}})

Comparing equation D.29 and equation D.30: in the former, the term pr(100c4101v)2>0p_{r}(\frac{100c_{4}}{101v})^{2}>0 is the extra constant (of scale Ω(1)\Omega(1), i.e. non-vanishing even under our asymptotic assumptions Assumption 3) compared with the latter.

In Theorem 3, we specify some necessary conditions that the optimal γ\gamma must satisfy. It is challenging to precisely characterize the optima (to within o(1)o(1) error), because doing so may require explicitly writing those smaller scale terms hidden (in ±o(1)\pm o(1)) by our asymptotic setting (Assumption 3). Those smaller scale terms, however, do not affect our analysis, because these ±o(1)\pm o(1) terms cannot reverse the Ω(1)\Omega(1) constant separation between the loss in the above different cases.

Our Stage-2 analysis on the optimal attention weights (equation 5) is based on freezing WV{\bm{W}}^{V} to be the Stage-1 optima characterized in Theorem 2. Notably, in Theorem 2, the uniqueness of the optima (i.e. a clean block-wise pattern) crucially depends on the L2L_{2} regularization. Indeed, as we prove in Theorem 4 (in Appendix C), without the regularization, there is a family of optima (depending on a series of free constants) all of which can encode the topic structure.

Among these alternative optima, we are particularly interested in a special case — one that has a diagonal pattern. This type of diagonally structured WV{\bm{W}}^{V} often occurs when we train the single-layered transformer model without L2L_{2} regularization.

Motivated by this empirical observation, we formally define the particular optima from Theorem 4 (in Appendix C) that is a diagonal pattern.

Corresponding to this case, we provide an analysis on the Stage-2 optimal attention weights, which shows a very interesting different behavior from the result in Theorem 3 (for block-wise WV{\bm{W}}^{V}).

Suppose the data distribution follows the topic modeling assumption in Section 3.1 and Assumption 1. Suppose we train a single layer transformer given by equation 7 with bpred=0{\bm{b}}^{\text{pred}}=0 and WV{\bm{W}}^{V} frozen to DV{\bm{D}}^{V} in Definition 2, under masked language modeling objective (equation 1) with the squared loss (equation 3), under Assumption 2, Assumption 3, and Assumption 4. Then, the optimal (α,β)(\alpha,\beta) satisfy:

Following the same steps leading to equation D.21, define:

in which q(x)≔11−(1−pc)pmx−pmpr(1−(1−pc)pm)Tvq(x)\coloneqq\frac{1}{1-(1-p_{c})p_{m}}x-\frac{p_{m}p_{r}}{(1-(1-p_{c})p_{m})Tv}

For a document w{\bm{w}} which contains topics t1,⋯ ,tτ∈[T]t_{1},\cdots,t_{\tau}\in[T], there are the following cases:

Recall that the label is X:j={1,i=wj0,i∈{0,⋯ ,Tv}\wj{\bm{X}}_{:j}=\begin{cases}1,\quad&i=w_{j}\\ 0,\quad&i\in\{0,\cdots,Tv\}\backslash w_{j}\end{cases}

Plugging in the asymptotics from Assumption 3, the above becomes

Plugging in the asymptotics from Assumption 3, the above terms vanish.

Plugging in the asymptotics from Assumption 3, the above terms vanish.

Plugging in the asymptotics from Assumption 3, the above becomes

Adding equation D.33 and equation D.35, we can see in the asymptotic regime of interest:

in which the constant c4c_{4} is defined as

c4≔11−(1−pc)pm∈(1,2)c_{4}\coloneqq\frac{1}{1-(1-p_{c})p_{m}}\in(1,2) by Assumption 4.

We will again consider several possible cases for α,β\alpha,\beta in equation D.36.

Case 1, β≤1+c4pmpr100c4vτ\mathbf{\beta\leq\frac{1+c_{4}p_{m}p_{r}}{100c_{4}}v\tau}: then c4ββ+(v−1)α+(1+c4pmpr)vτ≤c4ββ+100c4β<1100\frac{c_{4}\beta}{\beta+(v-1)\alpha+(1+c_{4}p_{m}p_{r})v\tau}\leq\frac{c_{4}\beta}{\beta+100c_{4}\beta}<\frac{1}{100}, and hence

Case 2, β>1+c4pmpr100c4vτ\mathbf{\beta>\frac{1+c_{4}p_{m}p_{r}}{100c_{4}}v\tau}: we have the following subcases:

If α≥c4v−1β\alpha\geq\frac{c_{4}}{v-1}\beta, then c4ββ+(v−1)α+(1+c4pmpr)vτ<c4ββ+(v−1)α≤c4ββ+c4β<c41+c4\frac{c_{4}\beta}{\beta+(v-1)\alpha+(1+c_{4}p_{m}p_{r})v\tau}<\frac{c_{4}\beta}{\beta+(v-1)\alpha}\leq\frac{c_{4}\beta}{\beta+c_{4}\beta}<\frac{c_{4}}{1+c_{4}}, and hence by equation D.36 L(α,β)≥pc(1−c41+c4)2+pr±o(1)L(\alpha,\beta)\geq p_{c}(1-\frac{c_{4}}{1+c_{4}})^{2}+p_{r}\pm o(1).

If α<c4v−1β\alpha<\frac{c_{4}}{v-1}\beta, then c4ββ+(v−1)α+(1c4pmpr+1)vT>c4ββ+c4β+(1c4pmpr+1)vT\frac{c_{4}\beta}{\beta+(v-1)\alpha+(\frac{1}{c_{4}p_{m}p_{r}}+1)vT}>\frac{c_{4}\beta}{\beta+c_{4}\beta+(\frac{1}{c_{4}p_{m}p_{r}}+1)vT}

If β≥c7(1c4pmpr+1)vT\beta\geq c_{7}(\frac{1}{c_{4}p_{m}p_{r}}+1)vT (for some constant c7≔1c4(v−1−1c4−1)c_{7}\coloneqq\frac{1}{c_{4}(\sqrt{v-1}-\frac{1}{c_{4}}-1)}), then c4ββ+(v−1)α+(1c4pmpr+1)vT>c4ββ+c4β+(1c4pmpr+1)vT≥c4ββ+c4β+1c7β=c41+c4+1c7\frac{c_{4}\beta}{\beta+(v-1)\alpha+(\frac{1}{c_{4}p_{m}p_{r}}+1)vT}>\frac{c_{4}\beta}{\beta+c_{4}\beta+(\frac{1}{c_{4}p_{m}p_{r}}+1)vT}\geq\frac{c_{4}\beta}{\beta+c_{4}\beta+\frac{1}{c_{7}}\beta}=\frac{c_{4}}{1+c_{4}+\frac{1}{c_{7}}}, and hence L(α,β)>pr[1+(c41+c4+1c7)2]±o(1)L(\alpha,\beta)>p_{r}[1+(\frac{c_{4}}{1+c_{4}+\frac{1}{c_{7}}})^{2}]\pm o(1)

Specifically: let α=τT\alpha=\sqrt{\tau T} and β=v−1c4−1α=v−1c4−1τT\beta=\frac{v-1}{c_{4}-1}\alpha=\frac{v-1}{c_{4}-1}\sqrt{\tau T}, then

Note that this is smaller than all previous cases, because

1v−1<(1−1100)2\frac{1}{v-1}<(1-\frac{1}{100})^{2} since vv is a large finite constant (see Assumption 3 and Assumption 4).

1v−1<(1−c41+c4)2\frac{1}{v-1}<(1-\frac{c_{4}}{1+c_{4}})^{2} since vv is a large finite constant (see Assumption 3 and Assumption 4).

1v−1<(c41+c4+1c7)2\frac{1}{v-1}<(\frac{c_{4}}{1+c_{4}+\frac{1}{c_{7}}})^{2} by the definition of c7c_{7} above.

Therefore, we conclude that all α,β>0\alpha,\beta>0 that minimize L(α,β)L(\alpha,\beta) must satisfy

D.5 Loss landscape with respect to attention weights in the non-asymptotic setting

When T,τT,\tau are finite, the loss expression turns out to be too complicated to characterize in closed form (because all the o(1)o(1) terms need to be expanded). So we instead numerically compute the loss landscape as a function of α\alpha and β\beta.

We set T=100T=100 following our experimental setup on Wikipedia dataset (in Section 6), and v=300v=300 (so total vocabulary size Tv=30000Tv=30000) following the pre-trained BERT tokenizer in Huggingface implementation Wolf et al. (2020). We will vary τ∈{20,40,60,80}\tau\in\{20,40,60,80\}.

First, when WV{\bm{W}}^{V} is fixed to a diagonal structure (Definition 2), Theorem 5 predicts that the loss is lowest when β\beta is within an interval (boundaries controlled by τ\tau and TT), and α\alpha is less than a constant multiple of β\beta. Both constraints are visible in the non-asymptotic setting, as we show in the following:

On the other hand, when WV{\bm{W}}^{V} is fixed to a block-wise structure with uniform blocks (i.e. optima in Theorem 2), Theorem 3 predicts that the loss is lowest when a convex combination of α\alpha and β\beta is within an interval (boundaries controlled by τ\tau and TT). As we show in the following, a variant of this constraint visibly holds in the non-asymptotic setting.

Appendix E ADDITIONAL EMPIRICAL RESULTS

In Theorem 2 and Figure 2 we have shown that when freezing uniform attention weights and one-hot word embedding, under L2L_{2}-regularization, training a single layer transformer on our synthetic topic modeling distribution (Section 3.1) would make its WV{\bm{W}}^{V} converge to a block-wise pattern that encodes the topic structure.

In the following Figure 9, we additionally show empirical results without L2L_{2}-regularization, matching our theory in Theorem 4.

Complementing our experimental results in Section 6, Figure 10 shows that even when the attention weights WK,WQ{\bm{W}}^{K},{\bm{W}}^{Q} are jointly trained with WV{\bm{W}}^{V}, the model would still approximately converge to the type of block-wise WV{\bm{W}}^{V} described in our analyses in Section 5.2.

E.2 Additional results on learned attention weights

Complementing our experimental results in Section 6, Table 3 shows that when the trained WV{\bm{W}}^{V} is closer to uniform within each block, i.e. on average, each word pays more attention to different words of the same topic than to words of different topics.

On the other hand, when the trained WV{\bm{W}}^{V} is closer to a diagonal pattern, the above ordering is partially reversed, Table 4 shows that on average, each word pays the most attention to the same word in the document, followed by words of different topics, and the least attention to different words of the same topic.

E.3 Additional details and results on natural language data

In particular, for fair comparison, we should focus on the embedding similarity and attention weights between different words of the same topic and different words of different topics. (This is because those metrics are less meaningful for a pair of two same words, since their embeddings dot product is expected to be larger, which further biases the attention score comparisons. )

We also note that, for each word, an LDA model assigns some probability distribution of its topics. To determine whether two words are of the same topic, it is more meaningful if they share a topic in which both words have high likelihood. (By contrast, if two words each has some rarely-used topic that happens to overlap, we intuitively think of them as having different topics.)

To formalize such intuition, we filter out stop tokens, and other tokens that are not central to any topic (determined by the LDA). That is, for each topic tt, LDA assigns to it a likelihood pip_{i} for each word wiw_{i} in the vocabulary (of size nn). We sort these (word, likelihood) pairs by decreasing likelihood:

then for a pre-defined threshold parameter θ∈(0,1)\theta\in(0,1) controlling the proportion of words to be assigned to each topic, we only consider the topic tt to contain the following words

Moreover, we note that sentence length may cause a bias in attention weights calculation: intuitively, the average attention weight is the inverse of sentence length, but longer sentences usually contain more topics (and hence a larger proportion of different-topic word pairs). Thus, we expect that the average attention weight between different-topic word pairs are smaller than that between same-topic word pairs, even for a transformer with random parameters. (Empirically this bias indeed exists robustly, both on synthetic data and on Wikipedia data.) Therefore, we debias the effect of sentence length on attention weights: for each sentence, while computing the pairwise attention weights among its words, we “normalize the sentence length to 100”, that is, we multiply the raw attention weights by sentence length, and then divide the result by 100. In this way, the average attention weight in each sentence is always 1100\frac{1}{100}, regardless of the proportion of same-topic and different-topic word pairs. Indeed, as Table 1 and Table 5 show, for a randomly initialized BERT model, after our debiasing, the average same-topic and different-topic attention weights are roughly equal.

For a set of pre-trained transformer-based models downloaded from Huggingface (Wolf et al., 2020), we compare the embedding similarity and attention weights between same-topic tokens and different-topic tokens. The topics are determined by fitting an LDA model with 100 topics on a sample of tokenized Wikipedia corpus. We apply the above-mentioned ambiguity filter and debiasing.

When we further restrict to keeping only one topic for each word (to be consistent with the setting in our theoretical analysis): see Table 1.

Without the last restriction above: see the following Table 5.