How Fine-Tuning Allows for Effective Meta-Learning

Kurtland Chua, Qi Lei, Jason D. Lee

Introduction

Meta-learning (Thrun & Pratt, 2012) has emerged as an essential tool for quickly adapting prior knowledge to a new task with limited data and computational power. In this context, a meta-learner has access to some different but related source tasks from a shared environment. The learner aims to uncover some inductive bias from the source tasks to reduce sample and computational complexity for learning a new task from the same environment. One of the most promising methods is representation learning (Bengio et al., 2013), i.e., learning a feature extractor (or a common representation) from the source tasks. At test time, a learner quickly adapts to the new task by fine-tuning the representation and retraining the final layer(s) (see, e.g., prototype networks (Snell et al., 2017)). Substantial improvement over directly learning from a single task is expected in few-shot learning (Antoniou et al., 2018), a setting that naturally arises in many applications including reinforcement learning (Mendonca et al., 2019; Finn et al., 2017), computer vision (Nichol et al., 2018), federated learning (McMahan et al., 2017) and robotics (Al-Shedivat et al., 2017).

The empirical success of representation learning has led to an increased interest in theoretical analyses of underlying phenomena. Recent work assumes an explicitly shared representation across tasks (Du et al., 2020; Tripuraneni et al., 2020b, a; Saunshi et al., 2020; Balcan et al., 2015). For instance, Du et al. (2020) shows a generalization risk bound consisting of an irreducible representation error and estimation error. Without fine-tuning the whole network, the representation error accumulates to the target task and is irreducible even with infinite target labeled samples. Due to the lack of (representation) fine-tuning during training, we refer to these methods as making use of “frozen representation” objectives. This result is consistent with empirical findings, which suggest substantial performance gains associated with fine-tuning the whole network, compared to just learning the final linear layer (Chen et al., 2020; Salman et al., 2020).

Requiring tasks to be linearly separable on the same features is also unrealistic for transferring knowledge to other domains (e.g., from ImageNet to medical images (Raghu et al., 2019)). Therefore, we consider a more realistic setting, where the available tasks only approximately share the same representation. We propose a theoretical framework for analyzing the sample complexity of fine-tuning using a representation derived from a MAML-like algorithm. We show that fine-tuning quickly adapts to new tasks, requiring fewer samples in certain cases compared to methods using “frozen representation” objectives (as studied in Du et al. (2020), and will be formalized in Section 2.2). To the best of our knowledge, no prior studies exist beyond fine-tuning a linear model (Denevi et al., 2018; Konobeev et al., 2020; Collins et al., 2020a; Lee et al., 2020) or only the task-specific layers (Du et al., 2020; Tripuraneni et al., 2020b, a; Mu et al., 2020). In particular, our work can be viewed as a continuation of the work presented in Tripuraneni et al. (2020b), where the authors have acknowledged that the framework does not incorporate representation fine-tuning, and thus is a promising line of future work.

The following outlines this paper and its contributions:

In Section 2, we outline the general setting and overall assumptions.

In Section 4, we extend the analysis to general function classes. Our result provides bounds of the form

In Section 5, we instantiate our guarantees in the logistic regression and two-layer neural network settings.

In Section 6, we extend the separation result presented in Section 3 to a non-linear setting.

In Section 7, we experimentally verify the separation result in the linear setting from Section 3.

The empirical success of MAML (Finn et al., 2017) or, more generally, meta-learning, invokes further theoretical analysis from both statistical and optimization perspectives. A flurry of work engages in developing more efficient and theoretically-sound optimization algorithms (Antoniou et al., 2018; Nichol et al., 2018; Li et al., 2017) or showing convergence analysis (Fallah et al., 2020; Zhou et al., 2019; Rajeswaran et al., 2019; Collins et al., 2020b). Inspired by MAML, a line of gradient-based meta-learning algorithms have been widely used in practice Nichol et al. (2018); Al-Shedivat et al. (2017); Jerfel et al. (2018). Much follow-up work focused on online-setting with regret bounds (Denevi et al., 2018; Finn et al., 2019; Khodak et al., 2019; Balcan et al., 2015; Alquier et al., 2017; Bullins et al., 2019; Pentina & Lampert, 2014).

The statistical analysis of meta-learning can be traced back to Baxter (2000); Maurer & Jaakkola (2005), with a principle of inductive bias learning. Following the same setting, Amit & Meir (2018); Konobeev et al. (2020); Maurer et al. (2016); Pentina & Lampert (2014) assume a shared meta distribution for sampling the source tasks and measure generalization error/gap averaged over the meta distribution. Another line of work connects the target performance to source data with some distance measure between distributions (Ben-David & Borbely, 2008; Ben-David et al., 2010; Mohri & Medina, 2012). Finally, another series of works studied the benefits of using additional “side information” provided with a task for specializing the parameters of an inner algorithm (Denevi et al., 2020, 2021).

The hardness of meta-learning also attracts some investigation under various settings. Some recent work studies the meta-learning performance in worst-case setting (Collins et al., 2020b; Hanneke & Kpotufe, 2020a, b; Kpotufe & Martinet, 2018; Lucas et al., 2020). Hanneke & Kpotufe (2020a) provides a no-free-lunch result with problem-independent minimax lower bound, and Konobeev et al. (2020) also provide problem-dependent lower bound on a simple linear setting.

General Setting

Let [n]≔{1,…,n}[n]\coloneqq\left\{1,\dots,n\right\}. We denote the vector L2L_{2}-norm as ∥⋅∥2\left\lVert\cdot\right\rVert_{2}, and the matrix Frobenius norm as ∥⋅∥F\left\lVert\cdot\right\rVert_{F}. Additionally, ⟨⋅,⋅⟩\left\langle\cdot,\cdot\right\rangle can denote either the Euclidean inner product or the Frobenius inner product between matrices.

For a matrix AA, we let σi(A)\sigma_{i}(A) denote its ithi^{\text{th}} largest singular value. Additionally, for positive semidefinite AA, we write λmax⁡(A)\lambda_{\max}(A) and λmin⁡(A)\lambda_{\min}(A) for its largest and smallest eigenvalues, and A1/2A^{1/2} for its principal square root. We write PAP_{A} for the projection onto the column span of AA, denoted Col⁡A\operatorname{Col}{A}, and PA⊥≔I−PAP_{A}^{\perp}\coloneqq I-P_{A} for the projection onto its complement.

We use standard O,ΘO,\Theta, and Ω\Omega notation to denote orders of growth. We also use a≲ba\lesssim b or a≪ba\ll b to indicate that a=O(b)a=O(b). Finally, we write a≍ba\asymp b for a=Θ(b)a=\Theta(b).

2 Problem Setting

We will refer to the procedure above as AdaptRep, as it can be intuitively described as finding an initialization ϕ\phi in representation space such that there exists a good representation nearby for every task (and thus a learner simply needs a small adaptation/fine-tuning step to achieve good performance). We note that the objective can be viewed as a constrained form of algorithms found in the literature such as iMAML (Rajeswaran et al., 2019) and Meta-MinibatchProx (Zhou et al., 2019). However, we do not have a train-validation split, as is widespread in practice. This setup is motivated by results in Bai et al. (2020), which show that data splitting may not be preferable performance-wise, assuming realizability. Furthermore, empirical evaluations have demonstrated successes despite the lack of such a split (Zhou et al., 2019).

In the following sections, we will analyze the performance of the learned predictor for the population loss:

We focus on the few-shot learning setting for the target task, where we assume limited access to target data, and a learner needs to effectively use the source tasks to learn the target task quickly.

For AdaptRep to be sensible, we need to ensure the existence of a desirable initialization. We do so by assuming that there exists an initialization θ0∗\theta_{0}^{\ast} such that for any t∈[T]t\in[T], there exists a representation θt∗\theta_{t}^{\ast} with ∥θt∗−θ0∗∥≤δ0\left\lVert\theta_{t}^{\ast}-\theta_{0}^{\ast}\right\rVert\leq\delta_{0} and a predictor wt∗w_{t}^{\ast} so that μt\mu_{t} is given by

Throughout the paper, we assume the existence of an oracle during source training time for solving (1), as in Du et al. (2020); Tripuraneni et al. (2020b). For detailed analyses of source training optimization, we refer the reader to Ji et al. (2020); Wang et al. (2020). Nevertheless, representation fine-tuning introduces nonconvexity during target time not present in prior work, where one only needed to solve the convex problem of optimizing a final linear layer. Thus, our bounds explicitly take into account optimization performance on (2). To this end, we analyze the use of projected gradient descent (PGD), which applies to a wide variety of settings, under certain loss landscape assumptions. These standalone results are also provided in Section LABEL:sec:pgd-performance-bound.

As a point of comparison with AdaptRep, we will also be analyzing the “frozen representation” objective used in Du et al. (2020); Tripuraneni et al. (2020b). In particular, using the notation introduced above, such objectives consider the following optimization problem:

Due to the fact that the representation, once chosen, is fixed/frozen for all source tasks, we refer to the representation learning method above as FrozenRep throughout the rest of the paper. In Sections 3.4 and 6, we will demonstrate that unlike with AdaptRep, there exists cases where FrozenRep is unable to take advantage of the fact that the tasks approximately share representations.

AdaptRep in the Linear Setting

This assumption is used in the proofs to guarantee probabilistic tail bounds, and can be replaced with other conditions with appropriate modifications to the analysis. Finally, we define q(⋅ | μ)∼N(μ,σ2)q\left(\cdot\ \middle|\ \mu\right)\sim\mathcal{N}\left(\mu,\sigma^{2}\right) for a fixed σ>0\sigma>0.

For any t∈[T]t\in[T], ∥wt∗∥2=Θ(1)\left\lVert w_{t}^{\ast}\right\rVert_{2}=\Theta(1), and σk2(W∗)=Ω(T/k)\sigma_{k}^{2}(W^{\ast})=\Omega(T/k).

Finally, we evaluate the performance of the learner on a target task θ∗≔B∗w∗+δ∗\theta^{\ast}\coloneqq B^{\ast}w^{\ast}+\delta^{\ast} for some w∗w^{\ast} and ∥δ∗∥2≤δ0\left\lVert\delta^{\ast}\right\rVert_{2}\leq\delta_{0}.

To gain an intuition for the parameters, note that (B∗+Δt∗)wt∗=B∗wt∗+δt∗(B^{\ast}+\Delta_{t}^{\ast})w_{t}^{\ast}=B^{\ast}w_{t}^{\ast}+\delta_{t}^{\ast}, where ∥δt∗∥2≲δ0\left\lVert\delta_{t}^{\ast}\right\rVert_{2}\lesssim\delta_{0} by the norm conditions in Assumption 3.2. Thus, we can think of the assumptions on the source predictor weights as ensuring their proximity to a rank-kk space.

Finally, as a convention since the parameterization is not unique, we define the optimal parameters so that (B∗)⊤Σδt∗=0(B^{\ast})^{\top}\Sigma\delta_{t}^{\ast}=0 for any t∈[T]t\in[T]. This results in no loss of generality, as we can always redefine wt∗w_{t}^{\ast} and δt∗\delta_{t}^{\ast} as such.

2 Training Procedure

(Source training) We can write the objective (1) in this setting as

However, note that the objective does not impose any constraint on the predictor, as we can offload the norm of Δt\Delta_{t} onto wtw_{t}. Therefore, we instead consider the regularized source training objective

As will be shown in Section A, the regularization is equivalent to regularizing λγ∥Δtwt∥2\sqrt{\lambda\gamma}\left\lVert\Delta_{t}w_{t}\right\rVert_{2}, which is consistent with the intuition that δt∗\delta_{t}^{\ast} has small norm.

(Target training) Letting B0B_{0} be the obtained representation after orthonormalization, we adapt to the target task by optimizing

as the feasible set, where we explicitly define c1c_{1} and c2c_{2} in Section A.

To understand the choice of training objective above, observe that the predictor can be written as

where the “antisymmetric” initialization scheme ensures that AB0w0=0A_{B_{0}}w_{0}=0. Due to the choice of Cβ\mathcal{C}_{\beta}, the first two terms are of norm O(1)O(1), while the last term is of norm O(1/β)O(1/\beta). Therefore, for large enough β\beta, we can treat the cross term Δw\Delta w as a negligible perturbation, and the predictor is approximately linear in the parameters. Indeed, one can show that Lβ\mathcal{L}_{\beta} is thus approximately convex, guaranteeing that the best solution found by PGD is nearly optimal.

3 Performance Bound

Now, we provide a performance bound on the performance of the algorithm proposed in the previous section, which we prove in Section A. We define the following ratesLog factors and non-dominant terms have been suppressed for clarity. Full rates are presented in the appendix. of interest:

4 A Hard Case for FrozenRep

In what follows, we demonstrate the existence of a family of task distributions satisfying the assumptions outlined in Section 3.1 that is difficult for the method in Du et al. (2020), which we will refer to as FrozenRepThis is in reference to the fact that the representation is frozen during source training, i.e. no task-specific fine-tuning.. More explicitly, we prove an Ω(d/n)\Omega(d/n) minimax rate on the target task when using an FrozenRep-derived representation, even with access to infinite source tasks and data. Since this rate is achievable via training on the target task directly, this demonstrates that FrozenRep fails to capture the shared information between the tasks. In contrast, by specializing Theorem 3.1 to the proposed family of tasks, we will show that AdaptRep can indeed achieve a strictly faster statistical rate.

In this setting, we can write the FrozenRep objective in (3) as

First, we characterize the span of B^\hat{B} in the limit of infinite source tasks and data. Intuitively, since both B∗wB^{\ast}w and δ\delta both lie in (distinct) rank-kk spaces, but (1/2ε)∥Σ1/2B∗w∥2≤∥Σ1/2δ∥2(1/\sqrt{2\varepsilon})\left\lVert\Sigma^{1/2}B^{\ast}w\right\rVert_{2}\leq\left\lVert\Sigma^{1/2}\delta\right\rVert_{2} for any (w,δ)(w,\delta) from the task distribution (and thus the error along δ\delta is larger), FrozenRep learns EkE_{k} rather than B∗B^{\ast}.

The claim is proven in Section LABEL:sec:erm-hard-case-proofs. Although “incorrect”, it is unclear a priori that this choice of B^\hat{B} is undesirable performance-wise – we now show that this indeed the case. In fact, any algorithm making use of B^\hat{B}, in the worst case, cannot perform any better than a learner that is constrained to only use target data.

Then, with high probability over the draw of nT≫dn_{T}\gg d samples during target training, we have that

The previous result, which we prove in Section LABEL:sec:erm-hard-case-proofs, shows that any procedure making use of target samples to learn a predictor of the form B^w^+δ^\hat{B}\hat{w}+\hat{\delta} has a minimax rate of Ω(d/n)\Omega(d/n). This includes target-time fine-tuning procedures used by methods in practice such as iMAML, MetaOptNet (Lee et al., 2019), and R2D2 (Bertinetto et al., 2019). Notably, this rate is achievable by performing linear regression solely on the target samples, reflecting that the FrozenRep learner failed to capture the shared task structure. In contrast, by specializing the guarantee of Theorem 3.1 to this setting, we have the following result:

AdaptRep in the Nonlinear Setting

We now describe a general framework for analyzing fine-tuning in general function classes. To simplify the notation, we modify the setting described in Section 2.2 so that both the representation ϕ\phi and task-specific weight vector ww are captured by one parameter θ∈Θ\theta\in\Theta, with corresponding predictor gθg_{\theta}.

Throughout this section, we denote the population loss induced by a predictor gg as

where μg\mu_{g} samples x∼px\sim p and y ∣ x∼q(⋅∣g(x))y\ |\ x\sim q(\cdot|g(x)). Additionally, we let Lg(h)\mathcal{L}^{g}(h) denote the corresponding finite-sample quantityNote that we have omitted the samples from the notation for brevity..

Furthermore, for a fixed parameter θ\theta and a set C\mathcal{C}, we define the set AθC\mathcal{A}^{\mathcal{C}}_{\theta} to be

Intuitively, we can think of C\mathcal{C} as the possible ways a learner could adapt, and AθC\mathcal{A}^{\mathcal{C}}_{\theta} is the resulting set of possible predictors given an initialization θ\theta. For convenience, we define (AθC)⊗T\left(\mathcal{A}_{\theta}^{\mathcal{C}}\right)^{\otimes T} to be the set of functions mapping XT→YT\mathcal{X}^{T}\to\mathcal{Y}^{T} defined as

That is, (AθC)⊗T\left(\mathcal{A}_{\theta}^{\mathcal{C}}\right)^{\otimes T} can be interpreted as the set of possible choices of TT predictors for TT source tasks that are all close to some initialization θ\theta.

2 Assumptions

We now outline the assumptions we make in this setting.

To ensure transfer from source to target, we impose the following condition, a specific instance of which was proposed by Du et al. (2020) in the linear setting, and proposed by Tripuraneni et al. (2020b) for general settings:

There exists constants (ν,ε)(\nu,\varepsilon) such that if ρ\rho is the distribution of target tasks, then for any θ∈Θ0\theta\in\Theta_{0},

That is, the average best-case performance on the set of target tasks is controlled by the task-averaged best-case performance on the source tasks.

The (ν,ε)(\nu,\varepsilon)-diversity assumption ensures that optimizing θ\theta for the average source task performance results in controlled average target task performance. Note that we weakened the condition in Tripuraneni et al. (2020b) to bound the average rather than worst-case target performance, as is more suitable for higher-dimensional settings.

3 Performance Bound

where (εij)(\varepsilon_{ij}) are i.i.d. Rademacher random variables and (xi)(x_{i}) are i.i.d. samples from some (preset) distribution.

Assume that Assumptions 4.1, 4.2, 4.3 and 4.4 all hold. Consider a learner following the training procedure outlined in Section 4.1, and let (θt)(\theta_{t}) be the set of iterates generated by PGD with appropriately chosen step size η\eta (specified in Section LABEL:sec:gen-perf-proofs). Then, with probability at least 1−δ1-\delta over the random draw of samples,

Note that the Rademacher complexity terms above decays in most settings as

Case Studies

As in the linear setting, we consider an input distribution pp with covariance Σ\Sigma. We restrict the set of labels Y\mathcal{Y} to {0,1}\left\{0,1\right\}, and consider the conditional distribution

where σ(y)=1/(1+e−y)\sigma(y)=1/(1+e^{-y}) is the sigmoid function.

There exists ρ>0\rho>0 such that if x∼ptx\sim p_{t}, then Σ−1/2x\Sigma^{-1/2}x is ρ2\rho^{2}-sub-Gaussian.

For any t∈[T]t\in[T], ∥wt∗∥2=Θ(1)≤r\left\lVert w_{t}^{\ast}\right\rVert_{2}=\Theta(1)\leq r, and σk2(W∗)=Ω(r2T/k)\sigma_{k}^{2}(W^{\ast})=\Omega(r^{2}T/k).

1.2 Training Procedure

1.3 Performance Guarantee

Having described the statistical assumptions and the training procedure, we now specialize the guarantee of Theorem 4.1 to this setting. Details are provided in Section LABEL:sec:case-study-proofs.

2 Two-Layer Neural Networks

For any t∈t\in, ∣σ′(t)∣≤L\left|\sigma^{\prime}(t)\right|\leq L and ∣σ′′(t)∣≤μ\left|\sigma^{\prime\prime}(t)\right|\leq\mu. Furthermore, σ(0)=0\sigma(0)=0.

Let θ0=(B0,w0)\theta_{0}=(B_{0},w_{0}) be an antisymmetric parameter. Then, for every xx, there exists feature vectors ϕB0(x)\phi_{B_{0}}(x) and ψB0,w0(x)\psi_{B_{0},w_{0}}(x) such that

We interpret the features ϕB0(x)\phi_{B_{0}}(x) and ψB0,w0(x)\psi_{B_{0},w_{0}}(x) to be the gradients of (B,w)↦f(B,w)β(x)(B,w)\mapsto f^{\beta}_{(B,w)}(x) evaluated at θ0\theta_{0}Closed-form expressions for these quantities are provided in Section LABEL:sec:case-study-proofs.. Additionally, ζB0,w0Δ,w(x)\zeta_{B_{0},w_{0}}^{\Delta,w}(x) is the Taylor error. Finally, we define ρB0,w0(x)\rho_{B_{0},w_{0}}(x) to be the concatenation of ϕB0(x)\phi_{B_{0}}(x) and ψB0,w0(x)\psi_{B_{0},w_{0}}(x). ∎

We will show that if ∥Δ∥F\left\lVert\Delta\right\rVert_{F} and ∥w∥F\left\lVert w\right\rVert_{F} are both O(1/β)O(1/\beta), then the remainder term ζB0,w0Δ,w(x)\zeta_{B_{0},w_{0}}^{\Delta,w}(x) is O(1/β)O(1/\beta), and thus the function class is approximately linear in (Δ,w)(\Delta,w) with feature functions ϕB0(⋅)\phi_{B_{0}}(\cdot) and ψB0,w0(⋅)\psi_{B_{0},w_{0}}(\cdot). Note that these features correspond to the “activation” and “gradient” features, respectively, that are empirically evaluated by Mu et al. (2020).

We now outline the statistical assumptions for this setting. For all tasks, the inputs are assumed to be sampled from a 11-norm-bounded distribution pp. Furthermore, we let q(⋅ ∣ μ)q(\cdot\ |\ \mu) be generated as μ+η\mu+\eta for some O(1)O(1)-bounded additive noise η\eta, similar to Tripuraneni et al. (2020b).

To define the source tasks, we fix a representation matrix B∗B^{\ast} and a linear predictor w0∗w_{0}^{\ast} so that θ0∗=(B∗,w0∗)\theta_{0}^{\ast}=(B^{\ast},w_{0}^{\ast}) is antisymmetric, as motivated by the previous discussion.

The assumption above is analogous to the diversity conditions assumed in the previous sections.

2.2 Training Procedure

2.3 Performance Guarantee

Having described the statistical assumptions and the training procedure, we proceed with the performance guarantee.

A FrozenRep Hard Case for the Nonlinear Setting

Following the discussion in Section 5.2, note that when we take β→∞\beta\to\infty, the resulting function class can be expressed as

as desired. As before, we consider the family of task distributions induced by any A∗A^{\ast} with Col⁡A∗⊆E∗\operatorname{Col}{A^{\ast}}\subseteq E^{\ast}.

With the above task distribution, we can then prove the following hardness result on FrozenRep:

Furthermore, let B^\hat{B} be the output of FrozenRep with access to infinite per-task samples and tasks, with the task distribution deterined by B∗B^{\ast}. Then, with high probability over the draw of nT≫dn_{T}\gg d samples during target training, we have that

In contrast, we have the following upper bound on the performance of AdaptRep:

Let B0=[A0,−A0]B_{0}=[A_{0},-A_{0}], and fix a θ∗∈SA∗\theta^{\ast}\in S_{A^{\ast}}. Consider a learner which solves

during target training, where B0B_{0} is the representation obtained from source training. Finally, we set k=Θ(1)k=\Theta(1) and ε=k/d\varepsilon=k/d. Then, with access to infinite per-task samples and tasks during source training, the learner achieves target loss bounded as

with probability at least 1−δ1-\delta over the draw of target samples.

Simulations

In this section, we experimentally verify the hard case for the linear setting presented in Section 3.4. Since the empirical success of MAML or its variants in general has already been demonstrated extensively in practice and in existing work, it is not the focus of this section.

During source training, both FrozenRep and AdaptRep are provided with 10001000 tasks and 10d10d samples per task from the task distribution in Section 3.4. During target time, we evaluate the learned representation on the worst-case regression task from the same family.

Before we detail our results, we briefly comment on the nonconvexity in the source training procedure. Rather than optimizing (4) or (7) during source training, we use an additional Frobenius-norm regularizer on B⊤B−WW⊤B^{\top}B-WW^{\top} to ensure that the two terms are balanced. In the case of FrozenRep, this regularized objective was shown to have a favorable optimization landscape in Tripuraneni et al. (2020a). We then used L-BFGS to optimize these regularized objectives. To further mitigate any possible optimization issues, we evaluated both methods with 1010 random restarts, and report the best of the 1010 restarts (as measured by the worst-case performance on the target task) for both methods.

(Subspace Alignment). First, we plot the alignment of the learned representation (using the best of the 1010 restarts described above) with the correct space B∗B^{\ast}. We measure this via the sine of the largest principal angle between the two spaces, i.e.

We plot the results in Figure 2. As predicted by Lemma 3.1, FrozenRep does not learn B∗B^{\ast}, in contrast to AdaptRep.

Conclusion

We have presented, to the best of our knowledge, the first statistical analysis of fine-tuning-based meta-learning. We demonstrate the success of such algorithms under the assumption of approximately shared representation between available tasks. In contrast, we show that methods analyzed by prior work that do not incorporate task-specific fine-tuning fail under this weaker assumption.

An interesting line of future work is to determine ways to formulate useful shared structure among MDPs, i.e. formulate settings for which meta-reinforcement learning succeeds and results in improved regret bounds for downstream tasks.

Acknowledgements

KC is supported by a National Science Foundation Graduate Research Fellowship, Grant DGE-2039656. QL is supported by NSF #2030859 and the Computing Research Association for the CIFellows Project. JDL acknowledges support of the ARO under MURI Award W911NF-11-1-0303, the Sloan Research Fellowship, and NSF CCF 2002272.

References

Appendix A Proof of Theorem 3.1

In this section, we will prove the performance guarantee in the linear representation setting presented in Theorem 3.1. We first compute a bound on the difference in the spans of the true underlying representation B∗B^{\ast} and the representation B0B_{0} obtained from training on the source tasks. Having done so, we then analyze the performance of the best predictor found by projected gradient descent.

For clarity of presentation, we will write θt∗=Bt∗wt∗+δt∗\theta_{t}^{\ast}=B_{t}^{\ast}w_{t}^{\ast}+\delta_{t}^{\ast} and θ^t=(B+Δt)wt\hat{\theta}_{t}=(B+\Delta_{t})w_{t} throughout this section. Furthermore, let δ^t=Δtwt\hat{\delta}_{t}=\Delta_{t}w_{t}. Finally, we will be making use of the following covariance concentration results throughout this section, allowing us to connect empirical averages to population averages and vice versa:

The proof is similar to that of Lemma A.1, and is omitted for brevity. ∎

Let B0B_{0} be a minimizer of (4), and {(Δt,wt)}t∈[T]\left\{(\Delta_{t},w_{t})\right\}_{t\in[T]} be minimizers for the inner optimization problem given B0B_{0}. If the regularizer coefficients are chosen such that

and γ=δ02λ\gamma=\delta_{0}^{2}\lambda, then with probability at least 1−δ/31-\delta/3,

Throughout this proof, we instantiate the high-probability event in Lemma A.1, which occurs with probability at least 1−δ/91-\delta/9.

Note that we can express δt∗\delta_{t}^{\ast} as [δt∗(wt∗)⊤/∥wt∗∥22]wt∗=Δt∗wt∗[\delta_{t}^{\ast}(w_{t}^{\ast})^{\top}/\left\lVert w_{t}^{\ast}\right\rVert_{2}^{2}]w_{t}^{\ast}=\Delta_{t}^{\ast}w_{t}^{\ast}. Thus, via the optimality of B0B_{0} and {(Δt,wt)}t∈[T]\left\{(\Delta_{t},w_{t})\right\}_{t\in[T]}, we can form the basic inequality

Note that the simplification of the regularizer on the optimum holds since ∥wt∗∥2=Θ(1)\left\lVert w_{t}^{\ast}\right\rVert_{2}=\Theta(1) by Assumption 3.2. Equivalently, by rearranging,

Finally, by Proposition LABEL:prop:sum-regularizers-equivalence, the regularizer on Δt\Delta_{t} and wtw_{t} can be rewritten as a regularizer on δ^t\hat{\delta}_{t}, i.e.

and observe that [θt∗−θ^t]t∈[T]∈S[\theta_{t}^{\ast}-\hat{\theta}_{t}]_{t\in[T]}\in S, by letting αt=B∗wt∗−B0wt\alpha_{t}=B^{\ast}w_{t}^{\ast}-B_{0}w_{t} and βt=δt∗−δ^t\beta_{t}=\delta_{t}^{\ast}-\hat{\delta}_{t}. We bound the right-hand side of (11) via bounding the supremum of the inner product over SS, i.e.

where S2={[βt]t∈[T] | ∥βt∥2≤δ0+∥δ^t∥2}S_{2}=\left\{[\beta_{t}]_{t\in[T]}\ \middle|\ \left\lVert\beta_{t}\right\rVert_{2}\leq\delta_{0}+\left\lVert\hat{\delta}_{t}\right\rVert_{2}\right\}. Note the abuse of notation in (I), where we say [αt]t∈[T]∈S[\alpha_{t}]_{t\in[T]}\in S if there exists [βt]t∈[T][\beta_{t}]_{t\in[T]} so that [αt+βt]t∈[T]∈S[\alpha_{t}+\beta_{t}]_{t\in[T]}\in S. This decomposes the Gaussian width into the sum of the Gaussian widths of a low-rank set (I) and a small norm set (II). We proceed to bound both terms accordingly to these two properties.

Bounding the Gaussian width of the low-rank set (I).

To bound the Gaussian width, we first enlarge SS to remove βt\beta_{t} from the definition of the feasible set. Fix any (αt,βt)(\alpha_{t},\beta_{t}) pair satisfying the conditions in SS, and note that

Therefore, by the reverse triangle inequality,

Consequently, we can enlarge the feasible set to

Having relaxed the constraints, we now proceed to the main argument.

We will bound (A), (B), and (C) individually.

For a fixed Vˉ\bar{V}, (A) is a chi-squared random variable with kTkT degrees of freedom scaled by σ2\sigma^{2}. However, since Vˉ\bar{V} depends on VV, we need to have a high probability bound for any element of the covering. By using known concentration bounds for chi-squared random variables together with the union-bound, we find that uniformly over the covering,

Taking the supremum over S1S_{1}, we thus obtain the following bound on the Gaussian width:

Note that the events for this sub-argument occur with probability at least 1−δ/91-\delta/9.

Bounding the Gaussian width of the low-norm set (II).

Furthermore, by the Hanson-Wright inequality, we have that with probability at least 1−δ/9T1-\delta/9T,

Putting everything together, we thus find that with probability at least 1−δ/91-\delta/9,

Having bounded both Gaussian widths, we can thus bound the right-hand side of (11) as

Finally, by solving the quadratic inequality using Proposition LABEL:prop:solve-quad-ineqs, we find that

from which the desired performance bound follows due to the concentration of the empirical covariance from Lemma A.1. Since the concentration of the source covariance matrices and each of the sub-arguments all hold with probability at least 1−δ/91-\delta/9, it follows that the main claim holds with probability at least 1−δ/31-\delta/3. ∎

Throughout this proof, we will write P≔PΣ1/2B0P\coloneqq P_{\Sigma^{1/2}B_{0}} and P⊥≔PΣ1/2B0⊥P^{\perp}\coloneqq P_{\Sigma^{1/2}B_{0}}^{\perp} for readability. To proceed, note that we intuitively expect that for a learner that has learned the correct spaces, PΣ1/2θ^t≈Σ1/2B∗wt∗P\Sigma^{1/2}\hat{\theta}_{t}\approx\Sigma^{1/2}B^{\ast}w_{t}^{\ast} as it is the low-rank component of the estimator, and consequently, P⊥Σ1/2θ^t≈Σ1/2δt∗P^{\perp}\Sigma^{1/2}\hat{\theta}_{t}\approx\Sigma^{1/2}\delta_{t}^{\ast}. Then, decomposing into the corresponding errors, we have that for any t∈[T]t\in[T],

We proceed to bound the inner product above. We do so by observing that if we were to replace θ^t\hat{\theta}_{t} by θt∗\theta_{t}^{\ast}, then

Putting everything together, we find that

This is a quadratic inequality in the two terms on the left-hand side, and so by applying Proposition LABEL:prop:solve-quad-ineqs,

where the last line follows by orthogonality. Finally, by Proposition LABEL:prop:mat-to-proj,

which together with the diversity assumption in Assumption 3.2 yields the final bound

A.2 Target Guarantees for the Linear Setting

Having established a connection between the performance on the source tasks and the difference in the spans of B∗B^{\ast} and B0B_{0}, we can now analyze the performance of the target training procedure. First, we bound the performance of nearly optimal points in Cβ\mathcal{C}_{\beta} for several possible choices of c1,c2c_{1},c_{2}.

Assume that (Δ,w)(\Delta,w) is ζ\zeta-suboptimal for Lβ\mathcal{L}_{\beta} with the constraint set Cβ={(Δ,w) | ∥Δ∥F≤c1/β,∥w∥2≤c2/β}\mathcal{C}_{\beta}=\left\{(\Delta,w)\ \middle|\ \left\lVert\Delta\right\rVert_{F}\leq c_{1}/\beta,\left\lVert w\right\rVert_{2}\leq c_{2}/\beta\right\}, i.e.

We write θ^\hat{\theta} for the predictor corresponding to (Δ,w)(\Delta,w), i.e. θ^=β(AB0+Δ)(w0+w)\hat{\theta}=\beta(A_{B_{0}}+\Delta)(w_{0}+w). Now, let

We proceed by proving the three cases separately. Throughout the proof, we instantiate the high-probability event in Lemma A.2, which guarantees that

Recall that this event occurs with probability at least 1−δ/91-\delta/9.

Due to the choice of c1c_{1} and c2c_{2}, there exists a parameter in Cβ\mathcal{C}_{\beta} corresponding to the prediction vector PXB0Xθ∗+Xδ∗P_{XB_{0}}X\theta^{\ast}+X\delta^{\ast}. Writing the corresponding basic inequality, we thus have that

Now, by the Hanson-Wright inequality, we can bound the last term as

with probability at least 1−δ/91-\delta/9. Therefore, we can rewrite the prior basic inequality as

Now, note that we can form the quadratic inequality

and thus by applying Proposition LABEL:prop:solve-quad-ineqs to solve the inequality and noting that PXB0z/σP_{XB_{0}}z/\sigma is distributed as a chi-squared random variable with kk degrees of freedom,

with probability at least 1−δ/91-\delta/9, which is the bound that we wanted to show.

c1=∥δˉ∥2,c2=∥wˉ∥2c_{1}=\left\lVert\bar{\delta}\right\rVert_{2},c_{2}=\left\lVert\bar{w}\right\rVert_{2}

Due to the choice of c1c_{1} and c2c_{2}, there exists a parameter in Cβ\mathcal{C}_{\beta} corresponding to a predictor that agrees with θ∗\theta^{\ast} on the target samples. Therefore,

We proceed with an argument similar to that used in the source guarantee, albeit simpler since the representation B0B_{0} is fixed (and thus no covering argument is required). Along these lines, we bound the first term using the low-rank of B0B_{0}. Via projections,

Therefore, by applying Proposition LABEL:prop:solve-quad-ineqs to solve the quadratic inequality, we have that with probability at least 1−δ/91-\delta/9,

To bound the second term, we simply make use of the norm constraints defining the feasible set Cβ\mathcal{C}_{\beta}, which we note is analogous to the low-norm sub-argument of the source guarantee. Formally, the Hanson-Wright inequality implies that with probability at least 1−δ/91-\delta/9,

Putting everything together, we thus have that

c1=∥θ∗∥2,c2=0c_{1}=\left\lVert\theta^{\ast}\right\rVert_{2},c_{2}=0

Due to the choice of c1c_{1} and c2c_{2}, there exists a parameter in Cβ\mathcal{C}_{\beta} corresponding to a predictor that agrees with θ∗\theta^{\ast} on the target samples. Therefore, we can write the basic inequality

Now, noting that PXz/σP_{X}z/\sigma is a chi-squared random variable with dd degrees of freedom, we have that with probability at least 1−2δ/91-2\delta/9,

Therefore, by solving the resulting quadratic inequality via Proposition LABEL:prop:solve-quad-ineqs, we obtain the bound

Observe that all relevant high-probability events for each case occur simultaneously with probability at least 1−δ/31-\delta/3, as desired. ∎

A.3 Optimization Landscape during Target Time Training

Having derived statistical rates on nearly-optimal points for several choices of Cβ\mathcal{C}_{\beta} in the prior section, all that remains to be shown is that projected gradient descent can indeed find such points. In particular, we will demonstrate that for large enough β\beta, the optimization landscape induced by Lβ\mathcal{L}_{\beta} is approximately convex. We do so by demonstrating that the objective satisfies the assumptions outlined in Section LABEL:sec:pgd-performance-bound, and thus the accompanying guarantees for projected gradient descent hold.

To bound the average squared Hessian operator norm, note that ∇θ2gθ(xi)[Δ,w]=βxi⊤Δw\nabla_{\theta}^{2}g_{\theta}(x_{i})[\Delta,w]=\beta x_{i}^{\top}\Delta w, which is independent of θ\theta. Then, by the variational characterization of the operator norm,

and thus ∥∇θ2gθ(xi)∥22≲β2∥xi∥22\left\lVert\nabla_{\theta}^{2}g_{\theta}(x_{i})\right\rVert_{2}^{2}\lesssim\beta^{2}\left\lVert x_{i}\right\rVert_{2}^{2}. Consequently,

where the first inequality uses the fact that ∥AB0⊤xi∥22=2∥B0⊤xi∥22=2∥PB0xi∥22≲∥xi∥22\left\lVert A^{\top}_{B_{0}}x_{i}\right\rVert_{2}^{2}=2\left\lVert B_{0}^{\top}x_{i}\right\rVert_{2}^{2}=2\left\lVert P_{B_{0}}x_{i}\right\rVert_{2}^{2}\lesssim\left\lVert x_{i}\right\rVert_{2}^{2}, by the definition of AB0A_{B_{0}} and assumed orthogonality of B0B_{0}. ∎

Throughout the proof, we will write λmax⁡\lambda_{\max} and λmin⁡\lambda_{\min} as shorthand for λmax⁡(Σ)\lambda_{\max}(\Sigma) and λmin⁡(Σ)\lambda_{\min}(\Sigma), respectively. By definition, since B0B_{0} is assumed to be orthonormal, the concentration of the sample covariance in Lemma A.2 implies that

Finally, we proceed to derive the final bound on δˉ\bar{\delta}. Note that Xδˉ=Xδ∗−PXB0Xδ∗+PXB0⊥XB∗w∗X\bar{\delta}=X\delta^{\ast}-P_{XB_{0}}X\delta^{\ast}+P_{XB_{0}}^{\perp}XB^{\ast}w^{\ast}. Therefore,

Now, by applying the properties of the trace operator and Proposition LABEL:prop:mat-to-proj,

A.4 Deducing Theorem 3.1

Having proven all the previous results, we can now assemble the main claim in Theorem 3.1. Recall that we have defined the rates