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 . We denote the vector -norm as , and the matrix Frobenius norm as . Additionally, can denote either the Euclidean inner product or the Frobenius inner product between matrices.
For a matrix , we let denote its largest singular value. Additionally, for positive semidefinite , we write and for its largest and smallest eigenvalues, and for its principal square root. We write for the projection onto the column span of , denoted , and for the projection onto its complement.
We use standard , and notation to denote orders of growth. We also use or to indicate that . Finally, we write for .
2 Problem Setting
We will refer to the procedure above as AdaptRep, as it can be intuitively described as finding an initialization 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 such that for any , there exists a representation with and a predictor so that 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 for a fixed .
For any , , and .
Finally, we evaluate the performance of the learner on a target task for some and .
To gain an intuition for the parameters, note that , where 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- space.
Finally, as a convention since the parameterization is not unique, we define the optimal parameters so that for any . This results in no loss of generality, as we can always redefine and 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 onto . Therefore, we instead consider the regularized source training objective
As will be shown in Section A, the regularization is equivalent to regularizing , which is consistent with the intuition that has small norm.
(Target training) Letting be the obtained representation after orthonormalization, we adapt to the target task by optimizing
as the feasible set, where we explicitly define and 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 . Due to the choice of , the first two terms are of norm , while the last term is of norm . Therefore, for large enough , we can treat the cross term as a negligible perturbation, and the predictor is approximately linear in the parameters. Indeed, one can show that 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 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 in the limit of infinite source tasks and data. Intuitively, since both and both lie in (distinct) rank- spaces, but for any from the task distribution (and thus the error along is larger), FrozenRep learns rather than .
The claim is proven in Section LABEL:sec:erm-hard-case-proofs. Although “incorrect”, it is unclear a priori that this choice of is undesirable performance-wise – we now show that this indeed the case. In fact, any algorithm making use of , 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 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 has a minimax rate of . 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 and task-specific weight vector are captured by one parameter , with corresponding predictor .
Throughout this section, we denote the population loss induced by a predictor as
where samples and . Additionally, we let denote the corresponding finite-sample quantityNote that we have omitted the samples from the notation for brevity..
Furthermore, for a fixed parameter and a set , we define the set to be
Intuitively, we can think of as the possible ways a learner could adapt, and is the resulting set of possible predictors given an initialization . For convenience, we define to be the set of functions mapping defined as
That is, can be interpreted as the set of possible choices of predictors for source tasks that are all close to some initialization .
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 such that if is the distribution of target tasks, then for any ,
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 -diversity assumption ensures that optimizing 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 are i.i.d. Rademacher random variables and 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 be the set of iterates generated by PGD with appropriately chosen step size (specified in Section LABEL:sec:gen-perf-proofs). Then, with probability at least 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 with covariance . We restrict the set of labels to , and consider the conditional distribution
where is the sigmoid function.
There exists such that if , then is -sub-Gaussian.
For any , , and .
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 , and . Furthermore, .
Let be an antisymmetric parameter. Then, for every , there exists feature vectors and such that
We interpret the features and to be the gradients of evaluated at Closed-form expressions for these quantities are provided in Section LABEL:sec:case-study-proofs.. Additionally, is the Taylor error. Finally, we define to be the concatenation of and . ∎
We will show that if and are both , then the remainder term is , and thus the function class is approximately linear in with feature functions and . 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 -norm-bounded distribution . Furthermore, we let be generated as for some -bounded additive noise , similar to Tripuraneni et al. (2020b).
To define the source tasks, we fix a representation matrix and a linear predictor so that 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 , the resulting function class can be expressed as
as desired. As before, we consider the family of task distributions induced by any with .
With the above task distribution, we can then prove the following hardness result on FrozenRep:
Furthermore, let be the output of FrozenRep with access to infinite per-task samples and tasks, with the task distribution deterined by . Then, with high probability over the draw of samples during target training, we have that
In contrast, we have the following upper bound on the performance of AdaptRep:
Let , and fix a . Consider a learner which solves
during target training, where is the representation obtained from source training. Finally, we set and . Then, with access to infinite per-task samples and tasks during source training, the learner achieves target loss bounded as
with probability at least 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 tasks and 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 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 random restarts, and report the best of the 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 restarts described above) with the correct space . 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 , 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 and the representation 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 and throughout this section. Furthermore, let . 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 be a minimizer of (4), and be minimizers for the inner optimization problem given . If the regularizer coefficients are chosen such that
and , then with probability at least ,
Throughout this proof, we instantiate the high-probability event in Lemma A.1, which occurs with probability at least .
Note that we can express as . Thus, via the optimality of and , we can form the basic inequality
Note that the simplification of the regularizer on the optimum holds since by Assumption 3.2. Equivalently, by rearranging,
Finally, by Proposition LABEL:prop:sum-regularizers-equivalence, the regularizer on and can be rewritten as a regularizer on , i.e.
and observe that , by letting and . We bound the right-hand side of (11) via bounding the supremum of the inner product over , i.e.
where . Note the abuse of notation in (I), where we say if there exists so that . 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 to remove from the definition of the feasible set. Fix any pair satisfying the conditions in , 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 , (A) is a chi-squared random variable with degrees of freedom scaled by . However, since depends on , 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 , we thus obtain the following bound on the Gaussian width:
Note that the events for this sub-argument occur with probability at least .
Bounding the Gaussian width of the low-norm set (II).
Furthermore, by the Hanson-Wright inequality, we have that with probability at least ,
Putting everything together, we thus find that with probability at least ,
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 , it follows that the main claim holds with probability at least . ∎
Throughout this proof, we will write and for readability. To proceed, note that we intuitively expect that for a learner that has learned the correct spaces, as it is the low-rank component of the estimator, and consequently, . Then, decomposing into the corresponding errors, we have that for any ,
We proceed to bound the inner product above. We do so by observing that if we were to replace by , 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 and , we can now analyze the performance of the target training procedure. First, we bound the performance of nearly optimal points in for several possible choices of .
Assume that is -suboptimal for with the constraint set , i.e.
We write for the predictor corresponding to , i.e. . 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 .
Due to the choice of and , there exists a parameter in corresponding to the prediction vector . 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 . 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 is distributed as a chi-squared random variable with degrees of freedom,
with probability at least , which is the bound that we wanted to show.
Due to the choice of and , there exists a parameter in corresponding to a predictor that agrees with on the target samples. Therefore,
We proceed with an argument similar to that used in the source guarantee, albeit simpler since the representation is fixed (and thus no covering argument is required). Along these lines, we bound the first term using the low-rank of . Via projections,
Therefore, by applying Proposition LABEL:prop:solve-quad-ineqs to solve the quadratic inequality, we have that with probability at least ,
To bound the second term, we simply make use of the norm constraints defining the feasible set , 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 ,
Putting everything together, we thus have that
Due to the choice of and , there exists a parameter in corresponding to a predictor that agrees with on the target samples. Therefore, we can write the basic inequality
Now, noting that is a chi-squared random variable with degrees of freedom, we have that with probability at least ,
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 , as desired. ∎
A.3 Optimization Landscape during Target Time Training
Having derived statistical rates on nearly-optimal points for several choices of 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 , the optimization landscape induced by 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 , which is independent of . Then, by the variational characterization of the operator norm,
and thus . Consequently,
where the first inequality uses the fact that , by the definition of and assumed orthogonality of . ∎
Throughout the proof, we will write and as shorthand for and , respectively. By definition, since 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 . Note that . 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