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 GG which outputs exact conditional token probabilities (Gi(x)=P[Xi ∣ x−i]G_{i}(x)=P[X_{i}\,|\,x_{-i}]) 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 X\mathcal{X} denote a finite vocabulary of input tokens, X∗\mathcal{X}^{*} the set of variable-length sequences of tokens, and X=(X1,…,XT)∈X∗X=(X_{1},\ldots,X_{T})\in\mathcal{X}^{*} a random sequence of TT tokens. Let Δ∣X∣\Delta^{|\mathcal{X}|} denote the space of probability distributions over tokens.

Let G(x)=(G1(x),G2(x),…)G(x)=(G_{1}(x),G_{2}(x),\ldots) denote the masked language model which predicts a probability vector for each timestep in the input xx. Our theoretical abstraction is that GiG_{i} perfectly computes the distribution of XiX_{i}, the ii-th token, conditioned on all other tokens: Gi(x)=P[Xi∣X−i=x−i]G_{i}(x)=P[X_{i}|X_{-i}=x_{-i}]. Here P[Xi ∣ X−i=x−i]∈Δ∣X∣P[X_{i}\,|\,X_{-i}=x_{-i}]\in\Delta^{|\mathcal{X}|} is a probability vector. In particular, Gi(x)G_{i}(x) does not depend on xix_{i}. The downstream task involves labeled examples (x,F⋆(x))∈X∗×Y(x,F^{\star}(x))\in\mathcal{X}^{*}\times{\mathcal{Y}}, where F⋆:X∗→YF^{\star}:\mathcal{X}^{*}\to{\mathcal{Y}} provides ground-truth downstream labels and Y{\mathcal{Y}} 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 H0H_{0} has the meaningful information for the downstream task, which is a binary classification task where the ground-truth labeling F⋆F^{\star} is assumed to be a linear classifier on the posterior P[H0 ∣ X1:T=x]P[H_{0}\,|\,X_{1:T}=x]:

The token emission probability matrix WW has linearly independent columns.

We also require the following regularity conditions on H0H_{0} and the state transitions.

The Markov chain H0,H1,…H_{0},H_{1},\ldots is ergodic, and P[H0]P[H_{0}] has full support.

We show that if WW has linearly independent columns, a linear head fits downstream labels.

where x′=(∅,x1:t)x^{\prime}=(\varnothing,x_{1:t}) is the concatenation of a special token ∅\varnothing with xx.We note that G1(x′)G_{1}(x^{\prime}) does not depend on x1′x_{1}^{\prime} and therefore x1′x_{1}^{\prime} can be any token.

The key for the proof is to leverage the following general statement about random variables U,V,ZU,V,Z such that U⊥V ∣ ZU\perp V\,|\,Z, which decomposes the expression for P[U ∣ V]P[U\,|\,V].

Let U,V,ZU,V,Z be random variables such that U⊥V ∣ ZU\perp V\,|\,Z. Then for any vv, P[U ∣ V=v]=P[U ∣ Z]⋅P[Z ∣ V=v]P[U\,|\,V=v]=P[U\,|\,Z]\cdot P[Z\,|\,V=v]. Thus, if P[U ∣ Z]P[U\,|\,Z] has a left inverse (P[U ∣ Z])†(P[U\,|\,Z])^{\dagger}, then P[Z ∣ V=v]=(P[U ∣ Z])†P[U ∣ V=v]P[Z\,|\,V=v]=(P[U\,|\,Z])^{\dagger}P[U\,|\,V=v].

By the conditional independence structure of the HMM, Proposition 3.4 immediately implies

where W†W^{\dagger} is the left inverse for WW, guaranteed to exist by Assumption 3.1. This lets us recover P[H1∣X2:T+1=x]P[H_{1}|X_{2:T+1}=x] by applying a linear function to G1(x′)G_{1}(x^{\prime}). Additional linear functions will be sufficient to obtain μ⊤P[H0∣X1:T=x]\mu^{\top}P[H_{0}|X_{1:T}=x] from P[H1∣X2:T+1=x]P[H_{1}|X_{2:T+1}=x]. 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 WW has full column rank, which implies the necessary condition that ∣H∣≤∣X∣|\mathcal{H}|\leq|\mathcal{X}|. Without this assumption, it is unclear how to recover P[H0 ∣ X1:T=x]P[H_{0}\,|\,X_{1:T}=x] from G(x)G(x) alone. However, in realistic settings we would expect ∣H∣>∣X∣|\mathcal{H}|>|\mathcal{X}|, 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 F⋆F^{\star} using a linear head on GG for HMMs where the non-degeneracy assumptions on WW 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 GG 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 GG to a model G‾\overline{G} that maps a sequence of embeddings e1,…,ete_{1},\dots,e_{t} to conditional probabilities G1(x),…,Gt(x)G_{1}(x),\dots,G_{t}(x) as follows. We observe that each token zz in the vocabulary X\mathcal{X} naturally corresponds to a ∣H∣|\mathcal{H}|-dimensional vector: the zz-th row of the emission probability matrix WW, or equivalently, P[Xi=z ∣ Hi]P[X_{i}=z\,|\,H_{i}]. We denote this embedding by e(z)e(z) and call the family of embeddings {e(z):z∈X}\{e(z):z\in\mathcal{X}\} proper embeddings. A fundamental property of HMMs is that the conditional probability P[Xi ∣ X−i=x−i]P[X_{i}\,|\,X_{-i}=x_{-i}] only depends on x1,…,xtx_{1},\dots,x_{t} through their embeddings e(x)=(e(x1),…,e(xt))e(x)=(e(x_{1}),\dots,e(x_{t})). In other words, there exists a function G‾i\overline{G}_{i} such that

In particular, we let G‾i\overline{G}_{i} compute the standard message passing algorithm that computes the conditional probability of HMMs. This ensures that G‾i\overline{G}_{i} is well defined on all sequences of nonnegative vectors in ∣H∣^{|\mathcal{H}|}, beyond sequences of proper embeddings.We assume that pretraining produces this G‾i\overline{G}_{i}, 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 u∈∣H∣u\in^{|\mathcal{H}|} to G‾\overline{G} in the first argument and proper embeddings at positions i>1i>1. We can interpret uu as the embedding of a fake token z~\widetilde{z}. Concretely, consider adding a new token z~\widetilde{z} to the vocabulary X\mathcal{X}, and changing the emission probability at position 1 to satisfy P[X1=z~ ∣ H1]=uP[X_{1}=\widetilde{z}\,|\,H_{1}]=u and for all z≠z~z\neq\widetilde{z}, P[X1=z ∣ H1]∝(1−u)⊙e(z)P[X_{1}=z\,|\,H_{1}]\propto(1-u)\odot e(z). Then G‾i(u,e(x1),…,e(xt))\overline{G}_{i}(u,e(x_{1}),\ldots,e(x_{t})) precisely computes the conditional probability P[Xi ∣ X−i=(z~,x1,…,xt)−i]P[X_{i}\,|\,X_{-i}=(\widetilde{z},x_{1},\dots,x_{t})_{-i}] under the modified HMM. We refer the readers to Section B for the formal definition of G‾i\overline{G}_{i} and formal proofs of the interpretation above.

There exists a set of essential hidden states H⋆⊆H\mathcal{H}^{\star}\subseteq\mathcal{H}, so that the columns of WW corresponding to H⋆\mathcal{H}^{\star}, {W:,h}h∈H⋆\{W_{:,h}\}_{h\in\mathcal{H}^{\star}} , are linearly independent. Furthermore, H⋆\mathcal{H}^{\star} covers all meaningful information for the downstream tasks: supp(μ)⊆H⋆\textup{supp}(\mu)\subseteq\mathcal{H}^{\star}.

In addition, a last technical requirement on H⋆\mathcal{H}^{\star} is as follows: there exists a set B⊆H{\mathcal{B}}\subseteq\mathcal{H} such that H⋆=∪h∈Bsupp(A:,h)\mathcal{H}^{\star}=\cup_{h\in{\mathcal{B}}}\textup{supp}(A_{:,h}). In other words, H⋆\mathcal{H}^{\star} must be the set of all states reachable by starting from some state in B{\mathcal{B}} and transitioning one step in the hidden Markov chain.

Compared to Assumption 3.1, which required that all columns of WW are linearly independent, Assumption 3.5 only requires linear independence on a subset H⋆\mathcal{H}^{\star} of essential states. In the setting where ∣H∣>∣X∣|\mathcal{H}|>|\mathcal{X}|, the condition for Theorem 3.3 can never hold. On the other hand, Assumption 3.5 could still hold, for example, if ∣supp(μ)∣<∣X∣|\textup{supp}(\mu)|<|\mathcal{X}| and the set of columns of WW corresponding to hidden states in supp(μ)\textup{supp}(\mu) is linearly independent. The last technical requirement in Assumption 3.5 is also required, which could be satisfied if columns of AA are sparse. The following theorem shows that when Assumption 3.5 holds, we can recover F⋆F^{\star} using soft prompt tuning with a linear head.

where e^\widehat{e} prepends uu 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 GG. This makes task-essential information easier to recover.

The key proof intuition is that although recovering P[H0 ∣ X1:T=x]P[H_{0}\,|\,X_{1:T}=x] is impossible without strong non-degeneracy conditions (Assumption 3.1), we can aim to recover P[H0 ∣ X1:T=x]P[H_{0}\,|\,X_{1:T}=x] on the subset of essential states H⋆\mathcal{H}^{\star} defined in Assumption 3.5, which suffices for computing μ⊤P[H0 ∣ X1:T=x]\mu^{\top}P[H_{0}\,|\,X_{1:T}=x], since H⋆⊇supp(μ)\mathcal{H}^{\star}\supseteq\textup{supp}(\mu). To recover P[H0 ∣ X1:T=x]P[H_{0}\,|\,X_{1:T}=x] on H⋆\mathcal{H}^{\star}, we observe in Lemma B.2 that prepending the prompt uu is equivalent to introducing a modified random sequence X^\widehat{X} and fake token z~\widetilde{z} which influences the posterior of H2H_{2} as follows:

for invertible diagonal matrix DD and positive scalar rxr_{x}. We choose uu such that the vector P[H2 ∣ X^1=z~]⊙P[H0 ∣ X1:T=x]P[H_{2}\,|\,\widehat{X}_{1}=\widetilde{z}]\odot P[H_{0}\,|\,X_{1:T}=x] is supported only on H⋆\mathcal{H}^{\star}. Because corresponding columns of WW are linearly independent by Assumption 3.5, we can then recover Pr(H0=h ∣ X1:T=x)\textup{Pr}(H_{0}=h\,|\,X_{1:T}=x) for h∈H⋆h\in\mathcal{H}^{\star} by applying a linear function to G‾2(e^(x))\overline{G}_{2}(\widehat{e}(x)). This suffices for computing μ⊤P[H0 ∣ X1:T=x]\mu^{\top}P[H_{0}\,|\,X_{1:T}=x]. 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 WW, 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 WW. 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 H0,H1,…H_{0},H_{1},\ldots, meant to model the evolution of syntax, and a persistent “memory” M=(M1,…,MN)M=(M_{1},\ldots,M_{N}) with NN total cells, where each MiM_{i} takes values in a finite set M{\mathcal{M}}. The full joint probability is as follows:

The hidden state is modified to explicitly consist of a disentangled cell index J∈[N]J\in[N] and syntax state S∈SS\in{\mathcal{S}}, such that Hi=(Ji,Si)H_{i}=(J_{i},S_{i}) and H=[N]×S\mathcal{H}=[N]\times{\mathcal{S}}. To sample the token at timestep ii given the hidden state Hi=(Ji,Si)H_{i}=(J_{i},S_{i}), we first use JiJ_{i} to index the memory MM, obtaining the random variable MJiM_{J_{i}}. XiX_{i} is then sampled according to some time-invariant probability depending on MJi,Ji,SiM_{J_{i}},J_{i},S_{i}:

We consider how this model may generate the sentence “The cow in the pasture rolled on the grass’ happily.” M1M_{1} could store the subject (“cow”), M2M_{2} the location (“pasture”), M3M_{3} the sentiment (“happily”), and SiS_{i} could determine part-of-speech. For timesteps where “cow” and “rolled” are emitted Ji=1J_{i}=1 because we emit information related to the sentence subject. Timesteps for “pasture” and “grass” would have Ji=2J_{i}=2.

Because our generative model disentangles HH and MM, we can relax the non-degeneracy assumption on the token emission probabilities WW, compared to Theorem 3.3. The relaxed assumption only requires the columns {W:,(m,h)}m∈M,h∈H⋆\{W_{:,(m,h)}\}_{m\in{\mathcal{M}},h\in\mathcal{H}^{\star}} to be linearly independent in a subset H⋆\mathcal{H}^{\star} of “recoverable” hidden states, whereas Assumption 3.1 required all columns to be linearly independent.

There exists a set of recoverable hidden states H⋆={j⋆}×S⋆\mathcal{H}^{\star}=\{j^{\star}\}\times{\mathcal{S}}^{\star}, such that the collection of token emission probabilities from M×H⋆{\mathcal{M}}\times\mathcal{H}^{\star}, {W:,(m,h)}m∈M,h∈H⋆\{W_{:,(m,h)}\}_{m\in{\mathcal{M}},h\in\mathcal{H}^{\star}}, is a linearly independent set of vectors.

Furthermore, the span of these vectors must be disjoint from the span of token emission probabilities from M×(H∖H⋆){\mathcal{M}}\times(\mathcal{H}\setminus\mathcal{H}^{\star}): span({W:,(m,h)}m∈M,h∈H⋆)∩span({W:,(m,h′)}m∈M,h∈H∖H⋆)={0∣X∣}\textup{span}(\{W_{:,(m,h)}\}_{m\in{\mathcal{M}},h\in\mathcal{H}^{\star}})\cap\textup{span}(\{W_{:,(m,h^{\prime})}\}_{m\in{\mathcal{M}},h\in\mathcal{H}\setminus\mathcal{H}^{\star}})=\{\mathbf{0}_{|\mathcal{X}|}\}.

Note that the non-degeneracy condition of Theorem 3.3 would require {W:,(m,h)}m∈M,h∈H\{W_{:,(m,h)}\}_{m\in{\mathcal{M}},h\in\mathcal{H}} to be linearly independent, whereas Assumption 4.2 only requires linear independence for h∈H⋆h\in\mathcal{H}^{\star}. The second condition states that H⋆\mathcal{H}^{\star} and H∖H⋆\mathcal{H}\setminus\mathcal{H}^{\star} 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 h=(ji,si)h=(j_{i},s_{i}) for which the set {W:,(m,h)}m∈M\{W_{:,(m,h)}\}_{m\in{\mathcal{M}}} of token emission probabilities is fundamentally not very diverse, and therefore not linearly independent. For example, if the syntax sis_{i} indicates “article”, i.e. words such as “a”, “an”, and “the”, the token emission probabilities would carry little information about MjiM_{j_{i}} because the choice of article does not depend much on semantics, so columns corresponding to si=“article”s_{i}=\textup{``article''} would not be linearly independent, violating Assumption 3.1. However, Assumption 4.2 allows us to avoid this issue by placing such hh in H∖H⋆\mathcal{H}\setminus\mathcal{H}^{\star}, a set of hidden states which we can ignore, and only including hidden states which carry a lot of information about MM in H⋆\mathcal{H}^{\star}. In Example 4.1, when Ji=2J_{i}=2 (location), Si=“noun”S_{i}=\textup{``noun''}, the position ii should convey a lot about the location (in this case, “pasture”), so it is more reasonable to assume that {W:,m,h}m∈M\{W_{:,m,h}\}_{m\in{\mathcal{M}}} is linearly independent for this hidden state.

Thus, our aim is to focus on recovering information for the downstream task from positions ii where Hi∈H⋆H_{i}\in\mathcal{H}^{\star}. Formally, we define the following set of input sequences containing positions ii where the posterior of HiH_{i} given x−ix_{-i} concentrates on H⋆\mathcal{H}^{\star}:

The following theorem shows that under Assumption 4.2, we can recover F⋆F^{\star} using the attention head described above, if x∈Rx\in{\mathcal{R}} is nonempty. Note that R{\mathcal{R}} is nonempty if the posterior of HiH_{i} concentrates on H⋆\mathcal{H}^{\star} for some ii. 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 R{\mathcal{R}} as in (4.3). Then there exist an attention head on G(x)G(x) and token embeddings e(xi)e(x_{i}) such that the following holds for any x∈Rx\in{\mathcal{R}}:

where the function Attn is in the form described in (4.2).

The idea is to use the attention mechanism to attend to positions ii where supp(P[Hi ∣ X−i=x−i])⊆H⋆\textup{supp}(P[H_{i}\,|\,X_{-i}=x_{-i}])\subseteq\mathcal{H}^{\star}. The intuition of Assumption 4.2 is that such positions are more informative for recovering the latent posteriors; indeed, from the outputs Gi(x)G_{i}(x) at such ii, the value function in the attention will be able to recover P[Mj⋆ ∣ X1:T=x]P[M_{j^{\star}}\,|\,X_{1:T}=x]. A full proof is provided in Section C.1.

2 Guarantees for prompt-tuning

Assumption 3.2 holds on the Markov chain H0,H1,…H_{0},H_{1},\ldots. Furthermore, P[H0]P[H_{0}] is the stationary distribution: P[H0]=AP[H0]P[H_{0}]=AP[H_{0}], where AA is the transition matrix.

As before, we assume sparsity of μ\mu and some non-degeneracy of WW, though the assumption is more relaxed and easier to state compared to the vanilla HMM setting.

Let M⋆≜supp(μ){\mathcal{M}}^{\star}\triangleq\textup{supp}(\mu) denote the set of non-zero coordinates in μ\mu. There exists a set of recoverable hidden states H⋆\mathcal{H}^{\star}, such that the collection of token emission probabilities from M⋆×H⋆{\mathcal{M}}^{\star}\times\mathcal{H}^{\star}, {W:,(m,h)}m∈M⋆,h∈H⋆\{W_{:,(m,h)}\}_{m\in{\mathcal{M}}^{\star},h\in\mathcal{H}^{\star}}, is linearly independent.

Furthermore, the span of these vectors must be disjoint from the span of token emission probabilities from M⋆×(H∖H⋆){\mathcal{M}}^{\star}\times(\mathcal{H}\setminus\mathcal{H}^{\star}): span({W:,(m,h)}m∈M⋆,h∈H⋆)∩span({W:,(m,h′)}m∈M⋆,h∈H∖H⋆)={0∣X∣}\textup{span}(\{W_{:,(m,h)}\}_{m\in{\mathcal{M}}^{\star},h\in\mathcal{H}^{\star}})\cap\textup{span}(\{W_{:,(m,h^{\prime})}\}_{m\in{\mathcal{M}}^{\star},h\in\mathcal{H}\setminus\mathcal{H}^{\star}})=\{\mathbf{0}_{|\mathcal{X}|}\}.

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 supp(μ)\textup{supp}(\mu), whereas Assumption 4.2 considers all m∈Mm\in{\mathcal{M}}. 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 AA. We now state our result for recovering F⋆F^{\star} 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 uu and attention head on G‾(e^(x))\overline{G}(\widehat{e}(x)) and the token embeddings which can compute the ground-truth F⋆(x)F^{\star}(x) for any x∈Rx\in{\mathcal{R}}, defined in (4.3):

where e^\widehat{e} 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 MM to concentrate on supp(μ)\textup{supp}(\mu). As a result, all irrelevant information to the task is removed from G‾i(e^(x))\overline{G}_{i}(\widehat{e}(x)), making it easier to recover the task-specific information about the posterior of MM. 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 GG 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 X1:TX_{1:T} of length 129, where the first token X1=[MASK]X_{1}=\texttt{[MASK]}. 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 μ⊤P[H1 ∣ X−1=x−1]\mu^{\top}P[H_{1}\,|\,X_{-1}=x_{-1}], where μ\mu 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 ∣H∣=4|\mathcal{H}|=4 and varying memory sizes ∣M∣∈{2,3,5,7}|{\mathcal{M}}|\in\{2,3,5,7\}. The downstream label is generated by computing μ⊤P[M ∣ X−1=x−1]\mu^{\top}P[M\,|\,X_{-1}=x_{-1}], where μ\mu denotes the ground-truth weights. Viewing the memory HMM as a HMM where the component on M{\mathcal{M}} 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 H1H_{1} 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 i≥1i\geq 1. Then there exists a diagonal matrix DD such that for all x∈supp(P[X])x\in\textup{supp}(P[X]),

First, we note that by Assumption 3.2, P[Hi]P[H_{i}] has full support. As a consequence, Pr(Xi+1:t+i=x)>0\textup{Pr}(X_{i+1:t+i}=x)>0. By Bayes’ rule,

Note that the vector P[Hi]P[H0]\frac{P[H_{i}]}{P[H_{0}]} has finite and positive entries. The same applies to the ratio rx≜Pr(X1:T=x)Pr(Xi+1:T+i=x)r_{x}\triangleq\frac{\textup{Pr}(X_{1:T}=x)}{\textup{Pr}(X_{i+1:T+i}=x)}. Thus, we get the desired statement. ∎

By definition, G1(x′)=P[X1 ∣ X2:T+1=x]G_{1}(x^{\prime})=P[X_{1}\,|\,X_{2:T+1}=x]. Therefore, our goal is to rewrite P[H0 ∣ X1:T=x]P[H_{0}\,|\,X_{1:T}=x] as a linear function of P[X1∣X2:T+1=x]P[X_{1}|X_{2:T+1}=x] (up to a scaling which won’t affect the linear head prediction). Concretely, we will show

for a scalar rx≥0r_{x}\geq 0. With this equation, taking b=μ⊤Bb=\mu^{\top}B will give the desired result.

First, observe that P[X1 ∣ X2:T+1=x]=WP[H1 ∣ X2:T+1=x]P[X_{1}\,|\,X_{2:T+1}=x]=WP[H_{1}\,|\,X_{2:T+1}=x] by Proposition 3.4. Next, we apply Claim A.1 to obtain an invertible matrix DD such that for all x∈supp(P[X])x\in\textup{supp}(P[X]), P[H1∣X2:T+1=x]=rxDP[H0∣X1:T=x]P[H_{1}|X_{2:T+1}=x]=r_{x}DP[H_{0}|X_{1:T}=x], where rx>0r_{x}>0 is a scalar.

If WW has full row rank, it has a left inverse W†W^{\dagger} with W†W=I∣H∣×∣H∣W^{\dagger}W=I_{|\mathcal{H}|\times|\mathcal{H}|}. Choosing b=μD−1W†b=\mu D^{-1}W^{\dagger}, 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 G‾\overline{G} described in Section 3.1. The model G‾\overline{G} takes a sequence of embedding vectors v=(v1,…,vt)v=(v_{1},\ldots,v_{t}) as input and implements message passing to compute a sequence of tt outputs. We first define left and right messages δ←i+1→i(v)\overleftarrow{\delta}_{i+1\to i}(v) and δ→i−1→i(v)\overrightarrow{\delta}_{i-1\to i}(v) for i∈[t]i\in[t], as follows:

Next, we define the aggregated message at timestep ii by

Note that if Assumption 3.2 holds about the Markov chain H0,H1,…H_{0},H_{1},\ldots, τi(v)\tau_{i}(v) is always well-defined because P[Hi]P[H_{i}] will have full support. Note that for the proper embeddings e(xi)=P[Xi=xi ∣ Hi]e(x_{i})=P[X_{i}=x_{i}\,|\,H_{i}], where for x=(x1,…,xt)x=(x_{1},\ldots,x_{t}), we use e(x)=(e(x1),…,e(xt))e(x)=(e(x_{1}),\ldots,e(x_{t})), we can check via classical results on message passing that

Finally, we let the model model G‾\overline{G} compute

There is an edge case where the demoninator is 0, i.e. ∥τi(v)∥1=0\|\tau_{i}(v)\|_{1}=0. To make the behavior of G‾\overline{G} well-defined, in this case we set G‾i(v)=0∣X∣\overline{G}_{i}(v)=\mathbf{0}_{|\mathcal{X}|}. We observe that if the input embedding are obtained by e(x)e(x), G‾i(v)\overline{G}_{i}(v) indeed computes the desired conditional probability vector for x∈supp(P[X])x\in\textup{supp}(P[X]):

First we formalize the observation that soft prompt tuning is equivalent to adding a fake token z~\widetilde{z} to the vocabulary with emission probabilities at timestep 1 given by uu, and letting G‾\overline{G} compute conditional probabilities for this new distribution over sequences.

In the setting of Theorem 3.6, fix any prompt vector u∈∣H∣u\in^{|\mathcal{H}|}. Define the random variable X^\widehat{X} with the same emission probabilities as XX for i>1i>1: P[X^i ∣ Hi]=P[Xi ∣ Hi]P[\widehat{X}_{i}\,|\,H_{i}]=P[X_{i}\,|\,H_{i}]. For timestep 1, we define the emission probabilities of X^1\widehat{X}_{1} as follows:

In the above equations, z~\widetilde{z} is a fake token added to the vocabulary at timestep 1. It follows that for any ii, defining τi\tau_{i} as in (B.1)

As a consequence, it follows that for i>1i>1 and any xx such that (z~,∅,x)−i∈supp(P[X^−i])(\widetilde{z},\varnothing,x)_{-i}\in\textup{supp}(P[\widehat{X}_{-i}]),

For any xx with (z~,∅,x)−i∉supp(P[X^−i])(\widetilde{z},\varnothing,x)_{-i}\notin\textup{supp}(P[\widehat{X}_{-i}]), G‾i(e^(x))=0\overline{G}_{i}(\widehat{e}(x))=\mathbf{0}.

Next, the following lemma disentangles the influences of the fake token z~\widetilde{z} and the input sequence on the posterior distribution of the hidden variable.

In the setting above, there exists an invertible diagonal matrix DD such that for all xx such that (z~,x)∈supp(P[X^−2])(\widetilde{z},x)\in\textup{supp}(P[\widehat{X}_{-2}]), the following equation holds:

We now complete the proof of Theorem 3.6.

Let B{\mathcal{B}} be the set defined in Assumption 3.5 and define uu such that uh=1u_{h}=1 if h∈Bh\in{\mathcal{B}} and uh=0u_{h}=0 otherwise. First, we restrict our focus to xx such that (z~,x)∈supp(P[X^−2])(\widetilde{z},x)\in\textup{supp}(P[\widehat{X}_{-2}]). For these xx, we can apply Lemma B.1 and Lemma B.2 in the manner described in the proof sketch. This gives G‾2(e^(x))=rxWDv\overline{G}_{2}(\widehat{e}(x))=r_{x}WDv for v≜(A(u⊙P[H1]))⊙P[H0 ∣ X1:T=x]v\triangleq(A(u\odot P[H_{1}]))\odot P[H_{0}\,|\,X_{1:T}=x]. By definition of B{\mathcal{B}}, we have supp(A(u⊙P[H1]))=H⋆\textup{supp}(A(u\odot P[H_{1}]))=\mathcal{H}^{\star}, so supp(Dv)⊆H⋆\textup{supp}(Dv)\subseteq\mathcal{H}^{\star}. Thus, there is a matrix W†^\widehat{W^{\dagger}} such that

The existence of W†^\widehat{W^{\dagger}} is due to the fact that {W:,h}h∈H⋆\{W_{:,h}\}_{h\in\mathcal{H}^{\star}} is a linearly independent set of vectors, and supp(Dv)⊆H⋆\textup{supp}(Dv)\subseteq\mathcal{H}^{\star} whenever xx satisfies (z~,x)∈supp(P[X^−2])(\widetilde{z},x)\in\textup{supp}(P[\widehat{X}_{-2}]). Next, we note that a matrix BB exists such that (BDv)h=Pr(H0=h ∣ X1:T=x)(BDv)_{h}=\textup{Pr}(H_{0}=h\,|\,X_{1:T}=x) for h∈H⋆h\in\mathcal{H}^{\star} and (BDv)h=0(BDv)_{h}=0 otherwise. This is because DD is invertible, and supp(A(u⊙P[H1]))=H⋆\textup{supp}(A(u\odot P[H_{1}]))=\mathcal{H}^{\star}, so we can recover P[H0 ∣ X1:T=x]P[H_{0}\,|\,X_{1:T}=x] on coordinates in H⋆\mathcal{H}^{\star} by applying another coordinate-wise scaling. It follows that we can set b=μ⊤BW†^b=\mu^{\top}B\widehat{W^{\dagger}}. With this choice of bb, we compute

where the last equality follows because supp(μ)⊆H⋆\textup{supp}(\mu)\subseteq\mathcal{H}^{\star}. This completes the case where (z~,x)∈supp(P[X^−2])(\widetilde{z},x)\in\textup{supp}(P[\widehat{X}_{-2}]).

Otherwise, for (z~,x)∉supp(P[X^−2])(\widetilde{z},x)\notin\textup{supp}(P[\widehat{X}_{-2}]), by the behavior of G‾\overline{G} in Lemma B.1, G‾2(e^(x))=0\overline{G}_{2}(\widehat{e}(x))=\mathbf{0}, so any linear head must output b⊤G‾2(e^(x))=0b^{\top}\overline{G}_{2}(\widehat{e}(x))=\mathbf{0}. Furthermore, by the conditional independence structure in X^\widehat{X}, we must also have supp(P[H2,X^1=z~])∩supp(P[H2,X^3:T+2=x])=∅\textup{supp}(P[H_{2},\widehat{X}_{1}=\widetilde{z}])\cap\textup{supp}(P[H_{2},\widehat{X}_{3:T+2}=x])=\emptyset. As supp(μ)⊆supp(P[H2,X^1=z~])\textup{supp}(\mu)\subseteq\textup{supp}(P[H_{2},\widehat{X}_{1}=\widetilde{z}]), this must also mean supp(μ)∩supp(P[H2,X^3:T+2=x])=∅\textup{supp}(\mu)\cap\textup{supp}(P[H_{2},\widehat{X}_{3:T+2}=x])=\emptyset. However, we also have P[H2,X^3:T+2=x]=P[H2,X3:T+2=x]P[H_{2},\widehat{X}_{3:T+2}=x]=P[H_{2},X_{3:T+2}=x] by the definition of X^\widehat{X}, and this must have the same support as P[H0 ∣ X1:T=x]P[H_{0}\,|\,X_{1:T}=x] by applying Claim A.1 and the fact that x∈supp(P[X])x\in\textup{supp}(P[X]). It follows that for this choice of xx, μ⊤P[H0 ∣ X1:T=x]=0\mu^{\top}P[H_{0}\,|\,X_{1:T}=x]=0, 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 τ\tau, and well-known results about message passing . Next, it suffices to consider the case where (z~,∅,x)−i∉supp(P[X^−i])(\widetilde{z},\varnothing,x)_{-i}\notin\textup{supp}(P[\widehat{X}_{-i}]), as the other case follows directly from the definition of G‾\overline{G} in terms of τ\tau. In this case, we observe that τi(e^(x))=P[Hi,X^−i=(z~,∅,x)−i]=0\tau_{i}(\widehat{e}(x))=P[H_{i},\widehat{X}_{-i}=(\widetilde{z},\varnothing,x)_{-i}]=\mathbf{0}. It follows that ∥τi(e^(x))∥1=0\|\tau_{i}(\widehat{e}(x))\|_{1}=0. Thus, from our definition of G‾\overline{G}, we must have G‾i(e^(x))=0\overline{G}_{i}(\widehat{e}(x))=\mathbf{0}. ∎

By the conditional independence relations in a HMM, X^1⊥X^3:T+2 ∣ H2\widehat{X}_{1}\perp\widehat{X}_{3:T+2}\,|\,H_{2}. Using Bayes’ rule, we obtain

Where we define rx≜Pr(X1:T=x)Pr(X^1=z~,X^3:T+2=x)r_{x}\triangleq\frac{\textup{Pr}(X_{1:T}=x)}{\textup{Pr}(\widehat{X}_{1}=\widetilde{z},\widehat{X}_{3:T+2}=x)}. We note that rxr_{x} is positive and well-defined by the conditions of the lemma and Theorem 3.6. We can set DD to be the matrix diag(1P[H0])\textup{diag}(\frac{\mathbf{1}}{P[H_{0}]}), 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 XiX_{i} is conditionally independent from everything else given MJi,Ji,SiM_{J_{i}},J_{i},S_{i}. However, we will prove this statement algebraically. We compute

Throughout this section, we use MJiM_{J_{i}} to denote the random variable obtained by indexing MM by JiJ_{i}, both of which are themselves random variables. Let I^\widehat{{\mathcal{I}}} denote the set of indices ii where supp(P[Ji ∣ X−i=x−i])={j⋆}\textup{supp}(P[J_{i}\,|\,X_{-i}=x_{-i}])=\{j^{\star}\} and supp(P[Si ∣ X−i=x−i])⊆S⋆\textup{supp}(P[S_{i}\,|\,X_{-i}=x_{-i}])\subseteq{\mathcal{S}}^{\star}. We will first construct the key function KK and query qq such that the set of I{\mathcal{I}} of attended-to positions (4.2) is precisely I^\widehat{{\mathcal{I}}}. This construction does not require the position embeddings β1,…,βt\beta_{1},\ldots,\beta_{t}, so we set them to 0\mathbf{0}.

The following lemma demonstrates the existence of KK and qq such that I=I^{\mathcal{I}}=\widehat{{\mathcal{I}}}.

The proof of Lemma C.2 requires the following claim.

In the last equality, we defined ν(h)\nu^{(h)} to be the expression in the parentheses. Note that ν(h)∈V(h)≜span({W:,(m,h)}m∈M)\nu^{(h)}\in{\mathcal{V}}^{(h)}\triangleq\textup{span}(\{W_{:,(m,h)}\}_{m\in{\mathcal{M}}}). Furthermore, for h∉H⋆h\notin\mathcal{H}^{\star}, ν(h)∈\widebarV≜span({W:,(m,h)}m∈M,h∈H∖H⋆)\nu^{(h)}\in\widebar{{\mathcal{V}}}\triangleq\textup{span}(\{W_{:,(m,h)}\}_{m\in{\mathcal{M}},h\in\mathcal{H}\setminus\mathcal{H}^{\star}}). As the spans (V(h))h∈H⋆({\mathcal{V}}^{(h)})_{h\in\mathcal{H}^{\star}} and \widebarV\widebar{{\mathcal{V}}} are all pairwise disjoint, by Assumption 4.2, for each h∈H⋆h\in\mathcal{H}^{\star}, we can recover

Now we have, for h∈H⋆h\in\mathcal{H}^{\star},

Likewise, the same reasoning gives 1⊤∑h∉H⋆ν(h)=∑h∉H⋆Pr(Hi=h ∣ X−i=x−i)1^{\top}\sum_{h\notin\mathcal{H}^{\star}}\nu^{(h)}=\sum_{h\notin\mathcal{H}^{\star}}\textup{Pr}(H_{i}=h\,|\,X_{-i}=x_{-i}). Thus, we can choose Θ(1)\Theta^{(1)} to be the matrix with rows Θh,:(1)=1⊤B(h)\Theta^{(1)}_{h,:}=\mathbf{1}^{\top}B^{(h)} when h∈H⋆h\in\mathcal{H}^{\star}, and for some arbitrary \widebarh∉H⋆\widebar{h}\notin\mathcal{H}^{\star}, Θ\widebarh,:(1)=1⊤\widebarB\Theta^{(1)}_{\widebar{h},:}=\mathbf{1}^{\top}\widebar{B}. We set all other rows to 0\mathbf{0}, and we can check that this satisfies the lemma requirements.

We now construct Θ(2,h)\Theta^{(2,h)}. We can express ν(h)\nu^{(h)} in a vectorized manner by writing

We choose the first ∣H∣|\mathcal{H}| entries of qq such that qh=1q_{h}=1 if h=(j⋆,s)h=(j^{\star},s) for s∈S⋆s\in{\mathcal{S}}^{\star}, and qh=0q_{h}=0 otherwise. The last entry is 0. Next, we choose Θ(K)\Theta^{(K)} so that the first ∣H∣|\mathcal{H}| rows are Θ(1)\Theta^{(1)}, and the last row is all zeros. where Θ(1)\Theta^{(1)} is defined in Claim C.3. With this choice of Θ(K)\Theta^{(K)}, K(Gi(x))h=Pr(Hi=h∣X−i=x−i)K(G_{i}(x))_{h}=\textup{Pr}(H_{i}=h|X_{-i}=x_{-i}) for h∈H⋆h\in\mathcal{H}^{\star}. Furthermore, ∥K(Gi(x))∥1=1\|K(G_{i}(x))\|_{1}=1, by Claim C.3.

Now we note that for all ii, 1=∥K(Gi(x))∥1≥q⊤K(Gi(x))1=\|K(G_{i}(x))\|_{1}\geq q^{\top}K(G_{i}(x)), and for i∈I^i\in\widehat{{\mathcal{I}}}, q⊤K(Gi(x))=∑s∈S⋆Pr(Hi=(j⋆,s)∣X−i=x−i)=1q^{\top}K(G_{i}(x))=\sum_{s\in{\mathcal{S}}^{\star}}\textup{Pr}(H_{i}=(j^{\star},s)|X_{-i}=x_{-i})=1 by definition of qq and I^\widehat{{\mathcal{I}}}. This implies that positions i∈I^i\in\widehat{{\mathcal{I}}} 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 i∈I^i\in\widehat{{\mathcal{I}}}.

We first choose Θ(V)\Theta^{(V)} such that the rows satisfy Θ(m,j⋆,s),:(V)=Θm,:(2,s)\Theta^{(V)}_{(m,j^{\star},s),:}=\Theta^{(2,s)}_{m,:} when s∈S⋆s\in{\mathcal{S}}^{\star} for Θ(2,s)\Theta^{(2,s)} constructed in Claim C.3, and Θ(m,j,s),:(V)=0∣X∣\Theta^{(V)}_{(m,j,s),:}=\mathbf{0}_{|\mathcal{X}|} otherwise for j≠j⋆j\neq j^{\star} or s∉S⋆s\notin{\mathcal{S}}^{\star}.

We claim that for i∈I^i\in\widehat{{\mathcal{I}}},

This is because for s∈S⋆s\in{\mathcal{S}}^{\star}, Θ(2,s)Gi(x)=P[Mj⋆,Hi=(j⋆,s) ∣ X−i=x−i]\Theta^{(2,s)}G_{i}(x)=P[M_{j^{\star}},H_{i}=(j^{\star},s)\,|\,X_{-i}=x_{-i}] by Claim C.3, and for h=(j,s)h=(j,s) for j≠j⋆j\neq j^{\star} or s∉S⋆s\notin{\mathcal{S}}^{\star},

Note that this last equality followed because Pr(Hi=h ∣ X−i=x−i)=0\textup{Pr}(H_{i}=h\,|\,X_{-i}=x_{-i})=0 for the choice of hh and i∈I^i\in\widehat{{\mathcal{I}}}. By construction of Θ(V)\Theta^{(V)}, these computations imply that (C.2) does indeed hold. The embedding can be chosen such that e(xi)=P[Xi=xi ∣ MJi,Ji,Si]e(x_{i})=P[X_{i}=x_{i}\,|\,M_{J_{i}},J_{i},S_{i}]. Thus, we have for i∈I^i\in\widehat{I}:

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 b=B⊤μb=B^{\top}\mu. It follows that for i∈I^i\in\widehat{{\mathcal{I}}},

We obtained the last equality by observing that ∑sP[Xi=xi,Mj⋆,Ji=j⋆,Si=s ∣ X−i=x−i]=P[Mj⋆,Xi=xi ∣ X−i=x−i]\sum_{s}P[X_{i}=x_{i},M_{j^{\star}},J_{i}=j^{\star},S_{i}=s\,|\,X_{-i}=x_{-i}]=P[M_{j^{\star}},X_{i}=x_{i}\,|\,X_{-i}=x_{-i}] for i∈I^i\in\widehat{{\mathcal{I}}}, as the distribution of HiH_{i} must concentrate where Ji=j⋆J_{i}=j^{\star}. Finally, we observe that μ⊤P[Mj⋆,Xi=xi ∣ X−i=x−i]=μ⊤P[Mj⋆ ∣ X1:T=x]Pr(Xi=xi ∣ X−i=x−i)\mu^{\top}P[M_{j^{\star}},X_{i}=x_{i}\,|\,X_{-i}=x_{-i}]=\mu^{\top}P[M_{j^{\star}}\,|\,X_{1:T}=x]\textup{Pr}(X_{i}=x_{i}\,|\,X_{-i}=x_{-i}), so setting rx,i=Pr(Xi=xi ∣ X−i=x−i)r_{x,i}=\textup{Pr}(X_{i}=x_{i}\,|\,X_{-i}=x_{-i}) 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 x∈supp(P[X])x\in\textup{supp}(P[X]) with I^\widehat{{\mathcal{I}}} (defined in Lemma C.2) nonempty, the attended-to positions I{\mathcal{I}} satisfy I=I^{\mathcal{I}}=\widehat{{\mathcal{I}}}, and V(Gi(x),e(xi))=rx,iμ⊤P[Mj⋆ ∣ X1:T=x]V(G_{i}(x),e(x_{i}))=r_{x,i}\mu^{\top}P[M_{j^{\star}}\,|\,X_{1:T}=x]. As the attention head computes the average of V(Gi(x),e(xi))V(G_{i}(x),e(x_{i})) over attended-to positions, and rx,ir_{x,i} is positive for all i∈I^i\in\widehat{{\mathcal{I}}}, 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 Ji=j⋆J_{i}=j^{\star} always, and we only need to consider the evolution of SiS_{i}.

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 H~0≜(M1,…,MN,H0),H~1≜(M1,…,MN,H1),...\widetilde{H}_{0}\triangleq(M_{1},\ldots,M_{N},H_{0}),\widetilde{H}_{1}\triangleq(M_{1},\ldots,M_{N},H_{1}),..., which uses the following transition probabilities:

We observe that η(P[Xi=xi ∣ MJi,(Ji,Si)])=P[Xi=xi ∣ H~i]\eta(P[X_{i}=x_{i}\,|\,M_{J_{i}},(J_{i},S_{i})])=P[X_{i}=x_{i}\,|\,\widetilde{H}_{i}].

We observe that this definition almost matches Section B, except it replaces HH with H~\widetilde{H}. Next, we define the aggregated message at timestep ii by

In the edge case where P[M]P[M] 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 τi\tau_{i} in this coordinate to 0. We will see that this preserves the meaning of the message τi\tau_{i}, which for the proper embeddings e(xi)=P[Xi=xi ∣ H~i]e(x_{i})=P[X_{i}=x_{i}\,|\,\widetilde{H}_{i}], with e(x)=(e(x1),…,e(xt))e(x)=(e(x_{1}),\ldots,e(x_{t})), computes

We observe that ϕ(τi(e(x)))=P[MJi,Ji,Si,X−i=x−i]∣M∣N−1\phi(\tau_{i}(e(x)))=\frac{P[M_{J_{i}},J_{i},S_{i},X_{-i}=x_{-i}]}{|{\mathcal{M}}|^{N-1}}. We now compute the model output as follows:

In the edge case where ∥ϕ(τi(v))∥1=0\|\phi(\tau_{i}(v))\|_{1}=0, we again define G‾(v)=0∣X∣\overline{G}(v)=\mathbf{0}_{|\mathcal{X}|}. We can observe that G‾i(e(x))=P[Xi ∣ X−i=x−i]\overline{G}_{i}(e(x))=P[X_{i}\,|\,X_{-i}=x_{-i}].

The downstream classifier uses the embedding e^(x)\widehat{e}(x) defined as follows:

The dimensions of the parameters b,Θ(V)b,\Theta^{(V)} 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 e^\widehat{e}:

where ϕ\phi is defined in (C.4). The following assumption extends Assumption 4.5 to the multiple memory case.

Let M⋆≜supp(μ){\mathcal{M}}^{\star}\triangleq\textup{supp}(\mu) denote the set of non-zero coordinates in μ\mu. There exists a set of recoverable hidden states H⋆\mathcal{H}^{\star}, such that the collection of token emission probabilities from M⋆×H⋆{\mathcal{M}}^{\star}\times\mathcal{H}^{\star}, {W:,(m,h)}m∈M⋆,h∈H⋆\{W_{:,(m,h)}\}_{m\in{\mathcal{M}}^{\star},h\in\mathcal{H}^{\star}}, is a linearly independent set of vectors.

Furthermore, define the following span of vectors:

Then \widebarV\widebar{{\mathcal{V}}} must be disjoint from the span of token emission probabilities from M⋆×H⋆{\mathcal{M}}^{\star}\times\mathcal{H}^{\star}:

Note that Assumption C.5 reduces to Assumption 4.5 the case where NN, 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 uu and attention head on G‾(e^(x))\overline{G}(\widehat{e}(x)) and the token embeddings which can compute the ground-truth F⋆(x)F^{\star}(x) for any x∈Rx\in{\mathcal{R}}, defined in (4.3):

Here e^\widehat{e} 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 z~\widetilde{z} 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 H~\widetilde{H} as in Section C.2. Fix any prompt vector u∈∣H~∣u\in^{|\widetilde{\mathcal{H}}|}. Define the random variable X^\widehat{X} with the same emission probabilities as XX for i>1i>1: P[X^i ∣ H~i]=P[Xi ∣ H~i]P[\widehat{X}_{i}\,|\,\widetilde{H}_{i}]=P[X_{i}\,|\,\widetilde{H}_{i}]. For timestep 1, we define the emission probabilities of X^1\widehat{X}_{1} as follows:

In the above equations, z~\widetilde{z} is a fake token added to the vocabulary at timestep 1. It follows that for any ii, defining τi\tau_{i} as in (C.3)

As a consequence, it follows that for i>1i>1 and any xx such that (z~,x)−i∈supp(P[X^−i])(\widetilde{z},x)_{-i}\in\textup{supp}(P[\widehat{X}_{-i}]),

For any ii and xx with (z~,x)−i∉supp(P[X^−i])(\widetilde{z},x)_{-i}\notin\textup{supp}(P[\widehat{X}_{-i}]), G‾i(e^(x))=0\overline{G}_{i}(\widehat{e}(x))=\mathbf{0}.

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 uu:

We will also use the notation x^≜(z~,x1,…,xt)\widehat{x}\triangleq(\widetilde{z},x_{1},\ldots,x_{t}). The following lemma considers behaviors in edge cases with this choice of uu.

Towards our proofs, the following result is useful.

In the setting of Theorem C.6, where P[H0]P[H_{0}] is the stationary distributions satisfying P[H0]=AP[H0]P[H_{0}]=AP[H_{0}], it holds that

Because P[H0]P[H_{0}] is stationary, we observe that P[M,Hi]=P[M,H0]P[M,H_{i}]=P[M,H_{0}] for all ii. We write

We will now restrict our focus to the set of inputs

Here S⋆{\mathcal{S}}^{\star} 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 I^\widehat{{\mathcal{I}}}, 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 Mj⋆M_{j^{\star}} has limited support.

In the setting of Theorem C.6 and Lemma C.7, let uu be defined as in (C.6). Then for all i>1i>1, supp(P[Mj⋆ ∣ X^−i=x^−i])⊆supp(μ)\textup{supp}(P[M_{j^{\star}}\,|\,\widehat{X}_{-i}=\widehat{x}_{-i}])\subseteq\textup{supp}(\mu) if Pr(X^−i=x^−i)>0\textup{Pr}(\widehat{X}_{-i}=\widehat{x}_{-i})>0.

In this equation we used -(1,i) to index all but the first and ii-th element of the sequence. We note that supp(P[X^1=z~ ∣ Mj⋆,M−j⋆=m−j⋆,H1=h])=supp(μ)\textup{supp}(P[\widehat{X}_{1}=\widetilde{z}\,|\,M_{j^{\star}},M_{-j^{\star}}=m_{-j^{\star}},H_{1}=h])=\textup{supp}(\mu) for all m−j⋆,hm_{-j^{\star}},h, 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 ν(h)\nu^{(h)} to be the expression in the parentheses. We consider several cases. First, when h=(j⋆,s)h=(j^{\star},s) for s∈Ss\in{\mathcal{S}}, we must have that when i>1i>1, P[Mj⋆ ∣ X^−i=x^−i]P[M_{j^{\star}}\,|\,\widehat{X}_{-i}=\widehat{x}_{-i}] is supported on M⋆{\mathcal{M}}^{\star} by Proposition C.11. Thus, ν(h)∈V(h)≜span({W:,(m,h)}m∈M⋆)\nu^{(h)}\in{\mathcal{V}}^{(h)}\triangleq\textup{span}(\{W_{:,(m,h)}\}_{m\in{\mathcal{M}}^{\star}}). As a result, for h∉H⋆h\notin\mathcal{H}^{\star}, ν(h)∈\widebarV\nu^{(h)}\in\widebar{{\mathcal{V}}}, which is the span of vectors defined in Assumption C.5. As the spans (V(h))h∈H⋆({\mathcal{V}}^{(h)})_{h\in\mathcal{H}^{\star}} and \widebarV\widebar{{\mathcal{V}}} are all pairwise disjoint, by Assumption 4.2, for each h∈H⋆h\in\mathcal{H}^{\star}, we can recover

The remainder of this proof for the construction of Θ(1)\Theta^{(1)} follows the same steps as Claim C.3.

We observe that because supp(P[Mj⋆,Hi=(j⋆,s) ∣ X^−i=x^−i])⊆M⋆\textup{supp}(P[M_{j^{\star}},H_{i}=(j^{\star},s)\,|\,\widehat{X}_{-i}=\widehat{x}_{-i}])\subseteq{\mathcal{M}}^{\star} by Proposition C.11, we can finish the proof by repeating the argument of Claim C.3. ∎

The following claim relating the support of HiH_{i} conditioned on X^\widehat{X} to the support of HiH_{i} conditioned on XX will also be useful.

In the setting of Theorem C.6 and Lemma C.7, suppose that uu is defined as in (C.6). For i>1i>1 with Pr(X^−i=x^−i)>0\textup{Pr}(\widehat{X}_{-i}=\widehat{x}_{-i})>0, we have

The last line used the time-invariance property of the HMM (Proposition C.8), the definition of x^\widehat{x}, and the fact that P[X^i ∣ Hi,M]P[\widehat{X}_{i}\,|\,H_{i},M] is distributed the same as P[Xi ∣ Hi,M]P[X_{i}\,|\,H_{i},M] for i>1i>1. On the other hand, note that P[Hi−1 ∣ X−(i−1)=x−(i−1)]=∑m,hP[M=m,H0=h,Hi−1 ∣ X−(i−1)=x−(i−1)]P[H_{i-1}\,|\,X_{-(i-1)}=x_{-(i-1)}]=\sum_{m,h}P[M=m,H_{0}=h,H_{i-1}\,|\,X_{-(i-1)}=x_{-(i-1)}]. 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 Θ(K)=[Θ(1)0]\Theta^{(K)}{}=\begin{bmatrix}\Theta^{(1)}\\ \mathbf{0}\end{bmatrix}, where Θ(1)\Theta^{(1)} is defined in Claim C.10, we obtain KK such that for all ii > 1, (K(G‾i(e^(x))))h=Pr(Hi=h∣X^−i=x^−i)(K(\overline{G}_{i}(\widehat{e}(x))))_{h}=\textup{Pr}(H_{i}=h|\widehat{X}_{-i}=\widehat{x}_{-i}) for h∈H⋆h\in\mathcal{H}^{\star}. Furthermore, (K(G‾i(e^(x))))∣H∣+1=0(K(\overline{G}_{i}(\widehat{e}(x))))_{|\mathcal{H}|+1}=0, and ∥K(G‾i(e^(x)))∥1=1\|K(\overline{G}_{i}(\widehat{e}(x)))\|_{1}=1. We choose β1=[0∣H∣−2]\beta_{1}=\begin{bmatrix}\mathbf{0}_{|\mathcal{H}|}\\ -2\end{bmatrix} and βi=0∣H∣+1\beta_{i}=\mathbf{0}_{|\mathcal{H}|+1} for i>1i>1. We also construct qq so that the first ∣H∣|\mathcal{H}| dimensions are the indicator on the set {j⋆}×S⋆\{j^{\star}\}\times{\mathcal{S}}^{\star}. We set q∣H∣+1=1q_{|\mathcal{H}|+1}=1. Note that this construction ensures that for i>1i>1, 1=∥K(G‾i(e^(x)))∥1≥q⊤(K(G‾i(e^(x)))+βi)≥01=\|K(\overline{G}_{i}(\widehat{e}(x)))\|_{1}\geq q^{\top}(K(\overline{G}_{i}(\widehat{e}(x)))+\beta_{i})\geq 0. Note that for i∈I^i\in\widehat{{\mathcal{I}}}, by Claim C.12 we have supp(P[Hi ∣ X^−i=x^−i])⊆supp(P[Hi−1 ∣ X−(i−1)=x−(i−1)])⊆{j⋆}×S⋆\textup{supp}(P[H_{i}\,|\,\widehat{X}_{-i}=\widehat{x}_{-i}])\subseteq\textup{supp}(P[H_{i-1}\,|\,X_{-(i-1)}=x_{-(i-1)}])\subseteq\{j^{\star}\}\times{\mathcal{S}}^{\star}. Thus, for such i∈I^i\in\widehat{{\mathcal{I}}}, we have q⊤(K(G‾i(e^(x)))+βi)=1q^{\top}(K(\overline{G}_{i}(\widehat{e}(x)))+\beta_{i})=1, achieving the maximum over all positions. Finally, we note that 1∉I1\notin{\mathcal{I}} because the position embedding β1\beta_{1} ensures that q⊤(K(G‾1(e^(x)))+β1)≤−1q^{\top}(K(\overline{G}_{1}(\widehat{e}(x)))+\beta_{1})\leq-1. Thus, I=I^{\mathcal{I}}=\widehat{{\mathcal{I}}}, as desired. ∎

Next, the following lemma constructs the value function, analogously to Lemma C.4.

As a consequence, for all i∈I^i\in\widehat{{\mathcal{I}}},

where rx,i>0r_{x,i}>0 is a positive scalar. In particular, this holds regardless of whether Pr(X^−i=x^−i)>0\textup{Pr}(\widehat{X}_{-i}=\widehat{x}_{-i})>0. Furthermore, when x^∉supp(P[X^])\widehat{x}\notin\textup{supp}(P[\widehat{X}]), for all i>1i>1, we must have

In the setting of Theorem C.6 and Lemma B.1 where uu takes the value in in (C.6), for all xx where x^≜(z~,x)∈supp(P[X^])\widehat{x}\triangleq(\widetilde{z},x)\in\textup{supp}(P[\widehat{X}]), we have

Now we have μ⊤diag(P[X^1=z~ ∣ Mj⋆,M−j⋆=m−j⋆,H1=h])=μ⊤\mu^{\top}\textup{diag}(P[\widehat{X}_{1}=\widetilde{z}\,|\,M_{j^{\star}},M_{-j^{\star}}=m_{-j^{\star}},H_{1}=h])=\mu^{\top} because by construction, P[X^1=z~ ∣ Mj⋆,M−j⋆=m−j⋆,H1=h]P[\widehat{X}_{1}=\widetilde{z}\,|\,M_{j^{\star}},M_{-j^{\star}}=m_{-j^{\star}},H_{1}=h] is only supported on supp(μ)\textup{supp}(\mu) 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 uu as in (C.6). Consider an input x∈supp(P[X])x\in\textup{supp}(P[X]) such that x^≜(z~,x1,…,xt)\widehat{x}\triangleq(\widetilde{z},x_{1},\ldots,x_{t}) satisfies Pr(X^=x^)=0\textup{Pr}(\widehat{X}=\widehat{x})=0. Then μ⊤P[Mj⋆∣X1:T=x]=0\mu^{\top}P[M_{j^{\star}}|X_{1:T}=x]=0. Furthermore, for any xx where Pr(X^−i=x^−i)=0\textup{Pr}(\widehat{X}_{-i}=\widehat{x}_{-i})=0 for some ii, we must have G‾i(e^(x))=0∣X∣\overline{G}_{i}(\widehat{e}(x))=\mathbf{0}_{|\mathcal{X}|}.

In particular, as supp(u)∩supp(P[M,H0,X1:T=x])=∅\textup{supp}(u)\cap\textup{supp}(P[M,H_{0},X_{1:T}=x])=\emptyset, it follows that Pr(Mj⋆=m,H0=h,X1:T=x)=0\textup{Pr}(M_{j^{\star}}=m,H_{0}=h,X_{1:T}=x)=0 for all m∈supp(μ)m\in\textup{supp}(\mu) and any hh, by the construction of uu. Since x∈supp(P[X])x\in\textup{supp}(P[X]), it follows that Pr(Mj⋆=m ∣ X1:T=x)=0\textup{Pr}(M_{j^{\star}}=m\,|\,X_{1:T}=x)=0 for all m∈supp(μ)m\in\textup{supp}(\mu), so μ⊤P[Mj⋆∣X1:T=x]=0\mu^{\top}P[M_{j^{\star}}|X_{1:T}=x]=0.

We note that the statement about G‾i(e^(x))\overline{G}_{i}(\widehat{e}(x)) follows because of Lemma C.7. ∎

To construct the value function, we define Θ(V)\Theta^{(V)} in the same manner as Lemma C.4, such that Θ(V)\Theta^{(V)} contains Θ(2,s)\Theta^{(2,s)} constructed in Claim C.10 as a submatrix: Θ(m,j⋆,s),:(V)=Θm,:(2,s)\Theta^{(V)}_{(m,j^{\star},s),:}=\Theta^{(2,s)}_{m,:} for s∈S⋆s\in{\mathcal{S}}^{\star}. All other rows of Θ(V)\Theta^{(V)} are 0\mathbf{0}. It now follows that for i∈I^i\in\widehat{{\mathcal{I}}} and xx where Pr(X^−i=x^−i)>0\textup{Pr}(\widehat{X}_{-i}=\widehat{x}_{-i})>0, by definition of I^\widehat{{\mathcal{I}}},

The proof that this claim is correct follows the same reasoning as Lemma C.4, where we argue that P[Hi ∣ X^−i=x^−i]P[H_{i}\,|\,\widehat{X}_{-i}=\widehat{x}_{-i}] must concentrate on {j⋆}×S⋆\{j^{\star}\}\times{\mathcal{S}}^{\star} for all i∈I^i\in\widehat{{\mathcal{I}}}. Thus, we can define b=B⊤μb=B^{\top}\mu, where BB is defined in Lemma C.4. We observe that for i∈I^i\in\widehat{{\mathcal{I}}}, the same reasoning as before gives

First, if (z~,x)∉supp(P[X^])(\widetilde{z},x)\notin\textup{supp}(P[\widehat{X}]), by Claim C.15, we have μ⊤P[Mj⋆ ∣ X1:T=x]=0\mu^{\top}P[M_{j^{\star}}\,|\,X_{1:T}=x]=0. The expression above must also equal , as (z~,x)∉supp(P[X^])(\widetilde{z},x)\notin\textup{supp}(P[\widehat{X}]). Otherwise, we have

Now we apply Claim C.14 to get the desired result in this case. A additional case is when Pr(X^−i=x^−i)=0\textup{Pr}(\widehat{X}_{-i}=\widehat{x}_{-i})=0. In this case, Claim C.15 shows that G‾i(e^(x))=0\overline{G}_{i}(\widehat{e}(x))=\mathbf{0}, so it follows that the value function also computes 0 in this case.

Finally, we need to check the case where x^∉supp(P[X^])\widehat{x}\notin\textup{supp}(P[\widehat{X}]), and we want to show V(G‾i(e^(x)),e^i(x))=0V(\overline{G}_{i}(\widehat{e}(x)),\widehat{e}_{i}(x))=0 for all i>1i>1. The case where Pr(X^−i=x^−i)=0\textup{Pr}(\widehat{X}_{-i}=\widehat{x}_{-i})=0 is already handled above. In the case where Pr(X^−i=x^−i)>0\textup{Pr}(\widehat{X}_{-i}=\widehat{x}_{-i})>0, we can apply Claim C.10 to our construction for Θ(V)\Theta^{(V)} to get

Thus, taking the element-wise product with ϕ(e(xi))=P[X^i=x^i ∣ MJi,Ji,Si]\phi(e(x_{i}))=P[\widehat{X}_{i}=\widehat{x}_{i}\,|\,M_{J_{i}},J_{i},S_{i}], we must have, by Proposition C.1,

Both of these terms must be 0 since x^∉supp(P[X^])\widehat{x}\notin\textup{supp}(P[\widehat{X}]), giving the desired result. ∎

The first case we consider is when x∈Zx\in{\mathcal{Z}}, 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 I^\widehat{{\mathcal{I}}} (C.8) is nonempty, the attended-to positions I{\mathcal{I}} satisfy I=I^{\mathcal{I}}=\widehat{{\mathcal{I}}}. In addition, by applying Lemma C.13, we also obtain that for x∈supp(P[X])x\in\textup{supp}(P[X]), V(G‾i(e^(x)),e^i(x))=rx,iμ⊤P[Mj⋆ ∣ X1:T=x]V(\overline{G}_{i}(\widehat{e}(x)),\widehat{e}_{i}(x))=r_{x,i}\mu^{\top}P[M_{j^{\star}}\,|\,X_{1:T}=x]. As the attention head averages V(G‾i(e^(x)),e^i(x))V(\overline{G}_{i}(\widehat{e}(x)),\widehat{e}_{i}(x)) over the attended-to positions, and rx,ir_{x,i} is positive for all i∈I^i\in\widehat{{\mathcal{I}}}, we obtain the desired result.

In the second case, x∉Zx\notin{\mathcal{Z}}, so (z~,x)∉supp(P[X^])(\widetilde{z},x)\notin\textup{supp}(P[\widehat{X}]). By Lemma C.13, for all i>1i>1, the value function outputs 0. However, by the construction in Lemma C.9, the attention will only attend to i>1i>1. Thus, the output of the attention head is . However, Claim C.15 also implies that μ⊤P[Mj⋆ ∣ X1:T=x]=0\mu^{\top}P[M_{j^{\star}}\,|\,X_{1:T}=x]=0, 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 μ\mu 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 μ\mu 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 [−0.5,0.5][-0.5,0.5]. Our implementation for prompt tuning used the code of , available at https://github.com/kipgparker/soft-prompt-tuning.