Why Do Pretrained Language Models Help in Downstream Tasks? An Analysis of Head and Prompt Tuning
Colin Wei, Sang Michael Xie, Tengyu Ma
Introduction
Natural language processing (NLP) has been revolutionized by large-scale pretrained language models such as BERT and GPT , which are adapted to a variety of downstream NLP tasks. Although a large body of empirical work seeks to understand the effectiveness of pretrained models , theoretical understanding is scarce. Theoretically analyzing the relationship between the pretraining and downstream tasks is challenging because pretraining and downstream settings can greatly differ.
The key starting point for our analysis is to link the pretraining and downstream settings through an underlying generative model of the data. We model the data distribution as a latent variable model and the downstream task as a function of the latent variables. Assuming that pretraining on a large corpus allows us to learn the generative model, the conditional token probabilities predicted by the pretrained model carry information about the hidden variables. In downstream adaptation, we aim to recover this information to solve the downstream task.
Though full finetuning is the de facto empirical standard, analyzing it is challenging because it requires characterizing the weights of the pretrained model. In this paper, we focus on head tuning and prompt tuning, which both freeze all pretrained parameters and allow us to treat the pretrained model as a black box. Head tuning trains task-specific heads on top of the pretrained model outputs. Prompt tuning optimizes a task-specific “prompt” that is concatenated to the model input. Studying prompt tuning is particularly interesting since it can match the performance of full finetuning with less computation time .
Our work contrasts with prior theoretical work , which assumes that downstream labels are recoverable via a linear head applied to the conditional token probabilities, and analyze how errors in pretraining or model misspecification propagate downstream. We consider specific generative distributions for which we can prove these assumptions, showing that head and prompt tuning can recover the downstream labels.
Our analysis considers two data-generating distributions with increasing realism. First, we consider data generated from a Hidden Markov Model (HMM), where the downstream task is to learn a linear classifier on the posterior distribution over the hidden states (Section 3). We prove that, under strong non-degeneracy conditions on token emission probabilities, a linear head applied to a pretrained model which outputs exact conditional token probabilities () can recover the downstream label (Theorem 3.3). Furthermore, we can prove better recovery guarantees with relaxed non-degeneracy assumptions (Assumption 3.1) by using continuous prompt tuning (Theorem 3.6), reflecting the strong empirical performance of prompt tuning . Intuitively, prompt tuning conditions the latent variables so that nonessential information for the downstream task can be ignored during the tuning phase, making task-essential information easier to recover.
Second, we also strengthen our analysis by leveraging additional structure in the data. Motivated by long-range dependences in natural language, we analyze HMM variants with additional latent “memory” variables that can store long-term information more easily than vanilla HMMs (Section 4). Here, the downstream task is to learn a linear classifier on the posterior distribution of the memory variables. We show that, under weaker non-degeneracy conditions than the first setting, an attention-based classification head can recover ground-truth downstream labels from pretrained model outputs (Theorem 4.3). Intuitively, our recovery guarantees improve because the classification head can focus on the persistent, task-essential information in the memory while ignoring other transient and nonessential aspects of the latent variables. As with the vanilla HMM, we analyze prompt tuning for relaxing the non-degeneracy conditions even further (Theorem 4.6).
In summary, we relate the pretraining and downstream tasks by assuming that the downstream task is to learn a classifier on the posterior distributions of the latent variables defined by an underlying generative model of text. Our theoretical contributions are: 1) in this setting we analyze an HMM generative model show that simple classification heads can recover the true downstream labels under certain non-degeneracy assumptions, 2) we prove that soft prompt tuning can relax the non-degeneracy assumptions needed for downstream recovery making it easier to extract task-specific information, and 3) our recovery guarantees are stronger for memory-augmented HMMs in comparison to the vanilla HMM when tuning an attention-based classfication head.
We empirically evaluate our theoretical results with language models pretrained on synthetically generated data from HMMs. We find that prompt tuning obtains good downstream performance when our non-degeneracy conditions are relaxed, whereas head tuning performs poorly. Furthermore, we show that head tuning obtains better downstream performance when data is generated from a memory-augmented HMM, compared to a vanilla HMM, as is predicted by our theory.Code is available at https://github.com/sangmichaelxie/pretraining_analysis.
The black box nature of BERT and related models has inspired a variety of empirical works which seek to understand them. Probing papers study whether a pretrained model computes various types of structured information (e.g., syntactic ) by evaluating the performance of simple classifiers, or probes, on the representations . Other papers ablate various aspects of pretraining, such as changing the masking scheme or permuting the word order .
In comparison, theoretical analysis of pretrained language models is limited. Besides , which we discussed in Section 1, Zhang and Hashimoto analyze using a linear classifier to approximately recover the latent variable in a Gaussian graphical model with sparse dependencies between observed variables. However, their analysis and setting are focused towards understanding syntactic dependencies between tokens, whereas we directly model and analyze downstream performance.
Prompt-based tuning , which has improved empirical downstream performance for lightweight adaptation methods beyond head tuning to approach full finetuning, is an important focus of our theoretical analysis. Shin et al. employ task-specific prompts that are optimized over the discrete token space. Schick and Schütze reformulate natural language tasks as cloze-style phrases to enable few-shot learning. Subsequent methods optimize “soft” prompts, or continuous embedding vectors. Lester et al. employ soft prompts on pretrained large-scale T5 models and show that as the model size increases, prompt tuning performance can eventually match finetuning. Hambardzumyan et al. applies a variant of soft prompt tuning to MLM models. Li and Liang propose prefix tuning, which prepends a trainable prefix embedding sequence to all layers of the transformer.
More broadly, Lee et al. analyze reconstruction-based self-supervised learning methods in a general setting and show that under certain conditional independence assumptions, predicting one observed variable from another allows recovery of the latent with a linear head. Other theoretical works analyzing self-supervised or constrastive learning include , but they are not directly relevant for our particular setting.
Formulations and notations
We analyze models pretrained on masked language modeling (MLM) objectives. Let denote a finite vocabulary of input tokens, the set of variable-length sequences of tokens, and a random sequence of tokens. Let denote the space of probability distributions over tokens.
Let denote the masked language model which predicts a probability vector for each timestep in the input . Our theoretical abstraction is that perfectly computes the distribution of , the -th token, conditioned on all other tokens: . Here is a probability vector. In particular, does not depend on . The downstream task involves labeled examples , where provides ground-truth downstream labels and is a discrete set of labels for classification.
Head and prompt tuning.
Analysis for Hidden Markov Models
Defining a relation between pretraining and downstream tasks is the foremost challenge for analysis. We propose to link the two via latent variable generative assumptions on the input distribution. We model the downstream task as a function of the posterior distribution of the latent variables. Towards a first result, this section studies the case where inputs are generated by HMMs (see Figure 1 (left)), which have been well-studied in the context of language and speech processing (see e.g. ).
Downstream tasks. We assume that has the meaningful information for the downstream task, which is a binary classification task where the ground-truth labeling is assumed to be a linear classifier on the posterior :
The token emission probability matrix has linearly independent columns.
We also require the following regularity conditions on and the state transitions.
The Markov chain is ergodic, and has full support.
We show that if has linearly independent columns, a linear head fits downstream labels.
where is the concatenation of a special token with .We note that does not depend on and therefore can be any token.
The key for the proof is to leverage the following general statement about random variables such that , which decomposes the expression for .
Let be random variables such that . Then for any , . Thus, if has a left inverse , then .
By the conditional independence structure of the HMM, Proposition 3.4 immediately implies
where is the left inverse for , guaranteed to exist by Assumption 3.1. This lets us recover by applying a linear function to . Additional linear functions will be sufficient to obtain from . We provide the full proof in Section A.
Proposition 3.4 is reminiscent of the arguments of , which leverages the independence structure in the same way. Subsequent sections will require more complicated analyses and recovery procedures.
A drawback of Theorem 3.3 is that it relies heavily on assuming has full column rank, which implies the necessary condition that . Without this assumption, it is unclear how to recover from alone. However, in realistic settings we would expect , as increasing the size of the hidden state space improves language modeling capabilities of HMMs .
In this section, we study applying soft, or continuous, prompt tuning to the setting above. We show that by using soft prompt tuning, we can recover using a linear head on for HMMs where the non-degeneracy assumptions on are relaxed. Our analysis provides insight into the empirical successes of prompt-tuning: intuitively, prompt tuning enables better recovery of the downstream task by conditioning the output of to only contain task-specific information.
Soft prompt tuning trains task-specific embedding vectors, but analyzing how the model processes embedding vectors is challenging because it requires opening up the black box of the pretrained model. Thus, we require additional abstractions about how the pretrained model processes the embedding vectors. We will extend the mask language model to a model that maps a sequence of embeddings to conditional probabilities as follows. We observe that each token in the vocabulary naturally corresponds to a -dimensional vector: the -th row of the emission probability matrix , or equivalently, . We denote this embedding by and call the family of embeddings proper embeddings. A fundamental property of HMMs is that the conditional probability only depends on through their embeddings . In other words, there exists a function such that
In particular, we let compute the standard message passing algorithm that computes the conditional probability of HMMs. This ensures that is well defined on all sequences of nonnegative vectors in , beyond sequences of proper embeddings.We assume that pretraining produces this , which we treat as a blackbox for prompt tuning.
In particular, for prompt tuning we can consider the case where we pass an arbitrary nonnegative vector to in the first argument and proper embeddings at positions . We can interpret as the embedding of a fake token . Concretely, consider adding a new token to the vocabulary , and changing the emission probability at position 1 to satisfy and for all , . Then precisely computes the conditional probability under the modified HMM. We refer the readers to Section B for the formal definition of and formal proofs of the interpretation above.
There exists a set of essential hidden states , so that the columns of corresponding to , , are linearly independent. Furthermore, covers all meaningful information for the downstream tasks: .
In addition, a last technical requirement on is as follows: there exists a set such that . In other words, must be the set of all states reachable by starting from some state in and transitioning one step in the hidden Markov chain.
Compared to Assumption 3.1, which required that all columns of are linearly independent, Assumption 3.5 only requires linear independence on a subset of essential states. In the setting where , the condition for Theorem 3.3 can never hold. On the other hand, Assumption 3.5 could still hold, for example, if and the set of columns of corresponding to hidden states in is linearly independent. The last technical requirement in Assumption 3.5 is also required, which could be satisfied if columns of are sparse. The following theorem shows that when Assumption 3.5 holds, we can recover using soft prompt tuning with a linear head.
where prepends to the input embedding sequence, as defined in (3.2).
Theorem 3.6 provides a stronger recovery result than Theorem 3.3, which only used a linear head. This is also reflected in our synthetic experiments (Section 5), and prior work which shows that variants of prompt tuning can perform much better than only training the last few layers of the model . Our theory suggests that prompt tuning could help by conditioning the hidden variables to remove nonessential information for the task from the output of . This makes task-essential information easier to recover.
The key proof intuition is that although recovering is impossible without strong non-degeneracy conditions (Assumption 3.1), we can aim to recover on the subset of essential states defined in Assumption 3.5, which suffices for computing , since . To recover on , we observe in Lemma B.2 that prepending the prompt is equivalent to introducing a modified random sequence and fake token which influences the posterior of as follows:
for invertible diagonal matrix and positive scalar . We choose such that the vector is supported only on . Because corresponding columns of are linearly independent by Assumption 3.5, we can then recover for by applying a linear function to . This suffices for computing . More details are in Section B.
Analysis for memory-augmented Hidden Markov Models
We study a memory-augmented HMM which explicitly disentangles the evolution of hidden states from a persistent “memory” variable. Inspired by natural sentences, this model is intended to better capture the distinction between syntax, which constantly evolves, and semantics, which changes less. This additional structure in the generative model allows us to strengthen our results by relaxing the non-degeneracy conditions on , the token emission probabilities. Thus, both head and prompt tuning are more powerful in this setting compared to Section 3 and can recover the downstream label with weaker non-degeneracy assumptions on . In Section 4.2, we show that soft prompt tuning also provides an advantage over head tuning alone.
Data distribution. The memory-augmented HMM, depicted in Figure 2, can be viewed as a generative variant of memory networks and is closely related to Hidden Topic Markov Models . There are two sets of latent variables in the memory-augmented HMM: a Markov chain on hidden states , meant to model the evolution of syntax, and a persistent “memory” with total cells, where each takes values in a finite set . The full joint probability is as follows:
The hidden state is modified to explicitly consist of a disentangled cell index and syntax state , such that and . To sample the token at timestep given the hidden state , we first use to index the memory , obtaining the random variable . is then sampled according to some time-invariant probability depending on :
We consider how this model may generate the sentence “The cow in the pasture rolled on the grass’ happily.” could store the subject (“cow”), the location (“pasture”), the sentiment (“happily”), and could determine part-of-speech. For timesteps where “cow” and “rolled” are emitted because we emit information related to the sentence subject. Timesteps for “pasture” and “grass” would have .
Because our generative model disentangles and , we can relax the non-degeneracy assumption on the token emission probabilities , compared to Theorem 3.3. The relaxed assumption only requires the columns to be linearly independent in a subset of “recoverable” hidden states, whereas Assumption 3.1 required all columns to be linearly independent.
There exists a set of recoverable hidden states , such that the collection of token emission probabilities from , , is a linearly independent set of vectors.
Furthermore, the span of these vectors must be disjoint from the span of token emission probabilities from : .
Note that the non-degeneracy condition of Theorem 3.3 would require to be linearly independent, whereas Assumption 4.2 only requires linear independence for . The second condition states that and are distinguishable by the token emission probabilities.
We explain Assumption 4.2 in the setting of Example 4.1. For natural language, there might be choices of for which the set of token emission probabilities is fundamentally not very diverse, and therefore not linearly independent. For example, if the syntax indicates “article”, i.e. words such as “a”, “an”, and “the”, the token emission probabilities would carry little information about because the choice of article does not depend much on semantics, so columns corresponding to would not be linearly independent, violating Assumption 3.1. However, Assumption 4.2 allows us to avoid this issue by placing such in , a set of hidden states which we can ignore, and only including hidden states which carry a lot of information about in . In Example 4.1, when (location), , the position should convey a lot about the location (in this case, “pasture”), so it is more reasonable to assume that is linearly independent for this hidden state.
Thus, our aim is to focus on recovering information for the downstream task from positions where . Formally, we define the following set of input sequences containing positions where the posterior of given concentrates on :
The following theorem shows that under Assumption 4.2, we can recover using the attention head described above, if is nonempty. Note that is nonempty if the posterior of concentrates on for some . For natural language, it is realistic to assume this can occur because syntactic aspects of a sentence are typically low-entropy when the full sentence is observed.
Assume that non-degeneracy (Assumption 4.2) and regularity (Assumption 3.2) hold. Define as in (4.3). Then there exist an attention head on and token embeddings such that the following holds for any :
where the function Attn is in the form described in (4.2).
The idea is to use the attention mechanism to attend to positions where . The intuition of Assumption 4.2 is that such positions are more informative for recovering the latent posteriors; indeed, from the outputs at such , the value function in the attention will be able to recover . A full proof is provided in Section C.1.
2 Guarantees for prompt-tuning
Assumption 3.2 holds on the Markov chain . Furthermore, is the stationary distribution: , where is the transition matrix.
As before, we assume sparsity of and some non-degeneracy of , though the assumption is more relaxed and easier to state compared to the vanilla HMM setting.
Let denote the set of non-zero coordinates in . There exists a set of recoverable hidden states , such that the collection of token emission probabilities from , , is linearly independent.
Furthermore, the span of these vectors must be disjoint from the span of token emission probabilities from : .
We note that Assumption 4.5, and Assumption C.5 for multiple memories, are relaxations of Assumption 4.2, as they only consider memory values in , whereas Assumption 4.2 considers all . An additional advantage of the memory-augmented HMM is that Assumption 4.2 is simpler than Assumption 3.1 and does not require any conditions on the transition matrix . We now state our result for recovering with soft prompt tuning and an attention head.
In the setting above, suppose that non-degeneracy Assumption 4.5 and stationarity Assumption 4.4 hold. Then there exists a prompt and attention head on and the token embeddings which can compute the ground-truth for any , defined in (4.3):
where is the embedding in (4.4) and Attn is defined in (4.2).
The intuition for this proof is similar to Theorem 3.6: the soft prompt conditions the memory to concentrate on . As a result, all irrelevant information to the task is removed from , making it easier to recover the task-specific information about the posterior of . A more general theorem statement for the multiple memories setting, and the full proof, is provided in Section C.3
Simulations
We empirically evaluate our theoretical results by pretraining a BERT-like masked language model (MLM) on synthetic data generated by an HMM. Our goal is to verify key implications of our theory in a more realistic setting where some assumptions, such as that outputs exact conditional probabilities, may not hold. First, we compare head and prompt tuning and show that prompt tuning improves downstream performance, especially when the recovery problem is degenerate. Second, we compare the effect of changing the data distribution from vanilla HMMs to memory-augmented HMMs on head tuning with an attention layer. We find that the downstream performance improves when the data has a long-term memory component. These observations support our theory. Our code is available at the following URL: https://github.com/sangmichaelxie/pretraining_analysis.
Pretraining data and downstream task. We generate pretraining data from an HMM with randomly generated transition matrix, emission probabilities, and start distributions. In all experiments, the HMMs have 10 vocabulary symbols, while the hidden state size varies. The downstream task uses input sequences of length 129, where the first token . We consider binary classifcation where labels are generated using linear functions of the analytically-computed posteriors in the HMMs. In all experiments, the ground truth linear weight is sparse with 6 nonzero entries at uniformly random locations with Gaussian values. More details are in Appendix D.
Head vs. prompt tuning. We compare head and prompt tuning as the hidden state size of the data-generating HMM varies. The downstream label is generated by computing , where is a random ground-truth linear weight. Head tuning learns a linear head on top of the softmax probabilities predicted by the pretrained model for filling in the first [MASK] token. Prompt tuning uses the same setup but also optimizes a length 20 continuous embedding and preprends it to the input sequence.
Figure 3 (left) shows that prompt tuning improves downstream performance substantially across all hidden state sizes ({4,8,10,15,25,30}). Prompt tuning improves especially when the hidden state size increases beyond the vocabulary size, which makes the recovery problem degenerate. Thus, as suggested by Theorem 3.6, prompt tuning helps relax the non-degeneracy conditions.
Memory-augmented HMMs. We investigate the effect of augmenting the data-generating HMM with a long-term memory. We consider the single memory case with and varying memory sizes . The downstream label is generated by computing , where denotes the ground-truth weights. Viewing the memory HMM as a HMM where the component on never changes, we can compare against the vanilla HMMs from the previous setting. For the memory-augmented HMM, we use head tuning with a single-cell attention layer on the entire sequence of softmax probability outputs. For the vanilla HMM in the comparison, we use a linear head on the output at the first position, as an attention head would perform worse since the downstream task depends only on and not any other timesteps.
Figure 3 (right) verifies that head tuning recovers the downstream task better when there is more structure in the data, as predicted by Theorem 4.3. Head tuning achieves near 100% downstream accuracy on all hidden state sizes.
Conclusion
We analyze how pretraining on generic language modeling tasks can improve performance on diverse downstream tasks. In our analysis framework, the downstream task requires predicting properties of the posterior distribution over latent variables in an underlying generative model. When the generative model is a standard HMM, downstream recovery is possible with a simple classification head under strong non-degeneracy assumptions. We also show that we can relax the non-degeneracy conditions by changing the generative model to a memory-augmented HMM or using prompt tuning. The generative distributions studied here are meant to provide a first-cut result – we also conjecture similar theorems to hold for other generative models, which we leave as an interesting direction for future work.
Another direction for future work is to analyze finetuning. Existing work analyzes finetuning for linear neural networks and obtains empirically useful insights , but analyzing neural networks with nonlinear activations is very challenging. Our analysis of head and prompt tuning treats the model as a black box. Analyzing finetuning requires understanding how to open up the black box, which is a major open question.
Acknowledgements
We thank Percy Liang, Tianyi Zhang, and Nelson Liu for helpful discussions. CW was supported by a NSF Graduate Research Fellowship. SMX was supported by a NDSEG Fellowship. TM acknowledges support of Google Faculty Award, NSF IIS 2045685, and JD.com.
References
Appendix A Proofs for Section 3
We provide the formal proof of Theorem 3.3 based on the sketch in Section 3. The following lemma will be useful in our analysis.
In the setting of Section 3, suppose that Assumption 3.2 holds. Fix any timestep . Then there exists a diagonal matrix such that for all ,
First, we note that by Assumption 3.2, has full support. As a consequence, . By Bayes’ rule,
Note that the vector has finite and positive entries. The same applies to the ratio . Thus, we get the desired statement. ∎
By definition, . Therefore, our goal is to rewrite as a linear function of (up to a scaling which won’t affect the linear head prediction). Concretely, we will show
for a scalar . With this equation, taking will give the desired result.
First, observe that by Proposition 3.4. Next, we apply Claim A.1 to obtain an invertible matrix such that for all , , where is a scalar.
If has full row rank, it has a left inverse with . Choosing , we obtain
Next, we complete the proof of Proposition 3.4.
Appendix B Formal abstraction for prompt tuning and proofs for Section 3.1
We first formalize the definition of the model described in Section 3.1. The model takes a sequence of embedding vectors as input and implements message passing to compute a sequence of outputs. We first define left and right messages and for , as follows:
Next, we define the aggregated message at timestep by
Note that if Assumption 3.2 holds about the Markov chain , is always well-defined because will have full support. Note that for the proper embeddings , where for , we use , we can check via classical results on message passing that
Finally, we let the model model compute
There is an edge case where the demoninator is 0, i.e. . To make the behavior of well-defined, in this case we set . We observe that if the input embedding are obtained by , indeed computes the desired conditional probability vector for :
First we formalize the observation that soft prompt tuning is equivalent to adding a fake token to the vocabulary with emission probabilities at timestep 1 given by , and letting compute conditional probabilities for this new distribution over sequences.
In the setting of Theorem 3.6, fix any prompt vector . Define the random variable with the same emission probabilities as for : . For timestep 1, we define the emission probabilities of as follows:
In the above equations, is a fake token added to the vocabulary at timestep 1. It follows that for any , defining as in (B.1)
As a consequence, it follows that for and any such that ,
For any with , .
Next, the following lemma disentangles the influences of the fake token and the input sequence on the posterior distribution of the hidden variable.
In the setting above, there exists an invertible diagonal matrix such that for all such that , the following equation holds:
We now complete the proof of Theorem 3.6.
Let be the set defined in Assumption 3.5 and define such that if and otherwise. First, we restrict our focus to such that . For these , we can apply Lemma B.1 and Lemma B.2 in the manner described in the proof sketch. This gives for . By definition of , we have , so . Thus, there is a matrix such that
The existence of is due to the fact that is a linearly independent set of vectors, and whenever satisfies . Next, we note that a matrix exists such that for and otherwise. This is because is invertible, and , so we can recover on coordinates in by applying another coordinate-wise scaling. It follows that we can set . With this choice of , we compute
where the last equality follows because . This completes the case where .
Otherwise, for , by the behavior of in Lemma B.1, , so any linear head must output . Furthermore, by the conditional independence structure in , we must also have . As , this must also mean . However, we also have by the definition of , and this must have the same support as by applying Claim A.1 and the fact that . It follows that for this choice of , , so the desired statement still stands. ∎
We fill in the proofs of the lemmas below.
First, we note that (B.2) follows directly from the derivation of , and well-known results about message passing . Next, it suffices to consider the case where , as the other case follows directly from the definition of in terms of . In this case, we observe that . It follows that . Thus, from our definition of , we must have . ∎
By the conditional independence relations in a HMM, . Using Bayes’ rule, we obtain
Where we define . We note that is positive and well-defined by the conditions of the lemma and Theorem 3.6. We can set to be the matrix , which has finite positive entries on the diagonal by Assumption 3.2. ∎
Appendix C Proofs for Section 4
First, we introduce a proposition which is generally useful for proving the theorems in Section 4.
In the setting of Section 4, it holds that
An alternative interpretation of this statement is that is conditionally independent from everything else given . However, we will prove this statement algebraically. We compute
Throughout this section, we use to denote the random variable obtained by indexing by , both of which are themselves random variables. Let denote the set of indices where and . We will first construct the key function and query such that the set of of attended-to positions (4.2) is precisely . This construction does not require the position embeddings , so we set them to .
The following lemma demonstrates the existence of and such that .
The proof of Lemma C.2 requires the following claim.
In the last equality, we defined to be the expression in the parentheses. Note that . Furthermore, for , . As the spans and are all pairwise disjoint, by Assumption 4.2, for each , we can recover
Now we have, for ,
Likewise, the same reasoning gives . Thus, we can choose to be the matrix with rows when , and for some arbitrary , . We set all other rows to , and we can check that this satisfies the lemma requirements.
We now construct . We can express in a vectorized manner by writing
We choose the first entries of such that if for , and otherwise. The last entry is 0. Next, we choose so that the first rows are , and the last row is all zeros. where is defined in Claim C.3. With this choice of , for . Furthermore, , by Claim C.3.
Now we note that for all , , and for , by definition of and . This implies that positions do indeed achieve the maximum attention scores. ∎
Next, we also require a construction of the value function such that it computes the correct prediction for all .
We first choose such that the rows satisfy when for constructed in Claim C.3, and otherwise for or .
We claim that for ,
This is because for , by Claim C.3, and for for or ,
Note that this last equality followed because for the choice of and . By construction of , these computations imply that (C.2) does indeed hold. The embedding can be chosen such that . Thus, we have for :
The last equality followed from applying the same reasoning as in Proposition C.1.
Now we pick the last linear weight in the value function by . It follows that for ,
We obtained the last equality by observing that for , as the distribution of must concentrate where . Finally, we observe that , so setting completes the proof. ∎
Now we can complete the proof of Theorem 4.3.
By applying Lemmas C.2 and C.4, we constructed key, query, and value functions for the attention head such that for all with (defined in Lemma C.2) nonempty, the attended-to positions satisfy , and . As the attention head computes the average of over attended-to positions, and is positive for all , we obtain the desired result. ∎
We note that this proof also works for the case where there is a single memory cell, as that is a special case where always, and we only need to consider the evolution of .
C.2 Formal abstraction for prompt tuning in Section 4.2
We will work directly in the case with multiple memories, as the single memory case is captured in this setting. We follow the construction in Section B. our message passing formulation requires the augmented Markov chain , which uses the following transition probabilities:
We observe that .
We observe that this definition almost matches Section B, except it replaces with . Next, we define the aggregated message at timestep by
In the edge case where does not have full support, the coordinate-wise division in the definition above would sometimes divide by 0. However, for all these cases both of the corresponding terms in the numerator must also be 0, so we can simply set the value of in this coordinate to 0. We will see that this preserves the meaning of the message , which for the proper embeddings , with , computes
We observe that . We now compute the model output as follows:
In the edge case where , we again define . We can observe that .
The downstream classifier uses the embedding defined as follows:
The dimensions of the parameters remain unchanged. Note that when there is just a single memory, this reduces to the case in Section 4.
C.3 Analysis for prompt tuning in the multiple memory setting
We will state and prove our result for the prompt tuning setting with multiple memories. For the multiple memory setting, the downstream classifier uses the following embedding function :
where is defined in (C.4). The following assumption extends Assumption 4.5 to the multiple memory case.
Let denote the set of non-zero coordinates in . There exists a set of recoverable hidden states , such that the collection of token emission probabilities from , , is a linearly independent set of vectors.
Furthermore, define the following span of vectors:
Then must be disjoint from the span of token emission probabilities from :
Note that Assumption C.5 reduces to Assumption 4.5 the case where , the number of memory cells, is 1. In any case, it is a relaxation of Assumption 4.2.
We now state and prove the result for multiple memories.
In the setting above, suppose that non-degeneracy Assumption C.5 and holds. In addition, suppose that Assumption 4.4 (stationarity) holds. Then there exists a prompt and attention head on and the token embeddings which can compute the ground-truth for any , defined in (4.3):
Here is the embedding in (4.4) and Attn is defined in (4.2).
We begin by rigorously stating the observation that soft prompt tuning is equivalent to adding a fake token to the vocabulary and modifying the token emission probabilities at timestep 1, analogous to Lemma B.1.
In the setting of Theorem C.6, define as in Section C.2. Fix any prompt vector . Define the random variable with the same emission probabilities as for : . For timestep 1, we define the emission probabilities of as follows:
In the above equations, is a fake token added to the vocabulary at timestep 1. It follows that for any , defining as in (C.3)
As a consequence, it follows that for and any such that ,
For any and with , .
The proof of Lemma C.7 mirrors the proof of Lemma B.1, so we omit it here.
In particular, throughout the proof we will use the following prompt :
We will also use the notation . The following lemma considers behaviors in edge cases with this choice of .
Towards our proofs, the following result is useful.
In the setting of Theorem C.6, where is the stationary distributions satisfying , it holds that
Because is stationary, we observe that for all . We write
We will now restrict our focus to the set of inputs
Here is defined in the non-degeneracy assumption. We will first construct key and query parameters such that the set of attended-to positions is precisely , following the proof of Theorem 4.3.
Towards proving Lemma C.9, the following construction will be useful.
Our proof will require the following result which shows that the distribution of has limited support.
In the setting of Theorem C.6 and Lemma C.7, let be defined as in (C.6). Then for all , if .
In this equation we used -(1,i) to index all but the first and -th element of the sequence. We note that for all , so the desired statement follows. ∎
The proof of this statement will be analogous to Claim C.3. As before, we have
In the last equality, we defined to be the expression in the parentheses. We consider several cases. First, when for , we must have that when , is supported on by Proposition C.11. Thus, . As a result, for , , which is the span of vectors defined in Assumption C.5. As the spans and are all pairwise disjoint, by Assumption 4.2, for each , we can recover
The remainder of this proof for the construction of follows the same steps as Claim C.3.
We observe that because by Proposition C.11, we can finish the proof by repeating the argument of Claim C.3. ∎
The following claim relating the support of conditioned on to the support of conditioned on will also be useful.
In the setting of Theorem C.6 and Lemma C.7, suppose that is defined as in (C.6). For with , we have
The last line used the time-invariance property of the HMM (Proposition C.8), the definition of , and the fact that is distributed the same as for . On the other hand, note that . This involves a sum over the same terms in the numerator in (C.9). Thus, as all the terms in the sum of (C.9) are nonnegative, the desired statement follows. ∎
This lets us complete the proof of Lemma C.9.
By setting , where is defined in Claim C.10, we obtain such that for all > 1, for . Furthermore, , and . We choose and for . We also construct so that the first dimensions are the indicator on the set . We set . Note that this construction ensures that for , . Note that for , by Claim C.12 we have . Thus, for such , we have , achieving the maximum over all positions. Finally, we note that because the position embedding ensures that . Thus, , as desired. ∎
Next, the following lemma constructs the value function, analogously to Lemma C.4.
As a consequence, for all ,
where is a positive scalar. In particular, this holds regardless of whether . Furthermore, when , for all , we must have
In the setting of Theorem C.6 and Lemma B.1 where takes the value in in (C.6), for all where , we have
Now we have because by construction, is only supported on and equals 1 on the support. Thus, we obtain
We also require the following result to handle edge cases where probability values are 0.
In the setting of Theorem C.6 and Lemma C.7, define as in (C.6). Consider an input such that satisfies . Then . Furthermore, for any where for some , we must have .
In particular, as , it follows that for all and any , by the construction of . Since , it follows that for all , so .
We note that the statement about follows because of Lemma C.7. ∎
To construct the value function, we define in the same manner as Lemma C.4, such that contains constructed in Claim C.10 as a submatrix: for . All other rows of are . It now follows that for and where , by definition of ,
The proof that this claim is correct follows the same reasoning as Lemma C.4, where we argue that must concentrate on for all . Thus, we can define , where is defined in Lemma C.4. We observe that for , the same reasoning as before gives
First, if , by Claim C.15, we have . The expression above must also equal , as . Otherwise, we have
Now we apply Claim C.14 to get the desired result in this case. A additional case is when . In this case, Claim C.15 shows that , so it follows that the value function also computes 0 in this case.
Finally, we need to check the case where , and we want to show for all . The case where is already handled above. In the case where , we can apply Claim C.10 to our construction for to get
Thus, taking the element-wise product with , we must have, by Proposition C.1,
Both of these terms must be 0 since , giving the desired result. ∎
The first case we consider is when , defined in (C.7). By applying Lemmas C.9 and C.13, we constructed key, query, and value functions for the attention head such that when (C.8) is nonempty, the attended-to positions satisfy . In addition, by applying Lemma C.13, we also obtain that for , . As the attention head averages over the attended-to positions, and is positive for all , we obtain the desired result.
In the second case, , so . By Lemma C.13, for all , the value function outputs 0. However, by the construction in Lemma C.9, the attention will only attend to . Thus, the output of the attention head is . However, Claim C.15 also implies that , giving the desired result. ∎
Appendix D Experimental details
For all experiments, we randomly generated the parameters of an HMM with 10 output symbols in its vocabulary. We generate a random transition matrix by taking a random convex combination of random permutation matrices. We mix as many permutation matrices as there are hidden states; i.e. if there are 4 hidden states, then we mix 4 random permutation matrices. The mixing weights are generated by sampling logits IID from a uniform distribution on $$ and then taking a softmax with temperature 0.01. Although this is a small temperature, the transition probabilities can still be around 0.7 for some transitions. The start distribution is also sampled in the same way, but with softmax temperature 10.0. The rows of the emission probability matrix is also sampled the same way with temperature 0.01.
Pretrain model.
The pretrained model follows the BERT-base architecture, except with 6 layers and a much smaller vocab size.
Pretrain data and task.
The pretraining data consists of 5000 sequences (documents) generated from the HMM, each with length 10240. We pretrain on this data by doing 5% masked LM on chunks of length 512. Pretraining runs for 3 epochs and takes about 5 hours on a single NVIDIA Tesla K80 GPU on 16-bit precision. We use an internal cluster for all experiments. Pretraining uses batch size 8 and learning rate 1e-5 with a linear warmup of 500 steps and linear decay schedule after 500 steps. We generated 20 pretraining (and downstream) datasets for each problem instance and average over the 20 runs in the vanilla HMM comparison, while the memory-based distributions are run for 5 trials of pretraining and finetuning.
Downstream.
The downstream task samples a sparse ground truth linear weight with 6 nonzero elements. Positions for nonzero entries are sampled uniformly at random and values are sampled i.i.d. from a standard normal distribution. Although we do binary classification, we sample with 2 rows and take the label to be the argmax of the two scores, instead of having 1 row and taking the sign. We find that this results in less degenerate datasets (datasets where all labels are the same).
We generate 5000 training, 500 validation and 1000 test examples for the downstream tasks. Downstream training uses learning rate 0.01 for both prompt tuning and head tuning, with a linear warmup/decay schedule, for 5 epochs over the downstream data. We take the model returned at the last checkpoint as the result (no early stopping). We found that it was important to train prompt tuning with full precision, since the gradients are relatively small and become zero with discretization.
We used message passing in the HMM to compute the posterior distributions of the latent variables analytically.
Prompt tuning.
We prepended a length 20 continuous prompt to each sequence of input word embeddings. We initialize elements of the prompt vectors IID from the uniform distribution on . Our implementation for prompt tuning used the code of , available at https://github.com/kipgparker/soft-prompt-tuning.