Few-Shot Learning via Learning the Representation, Provably
Simon S. Du, Wei Hu, Sham M. Kakade, Jason D. Lee, Qi Lei
Introduction
A popular scheme for few-shot learning, i.e., learning in a data-scarce environment, is representation learning, where one first learns a feature extractor, or representation, e.g., the last layer of a convolutional neural network, from different but related source tasks, and then uses a simple predictor (usually a linear function) on top of this representation in the target task. The hope is that the learned representation captures the common structure across tasks, which makes a linear predictor sufficient for the target task. If the learned representation is good enough, it is possible that a few samples are sufficient for learning the target task, which can be much smaller than the number of samples required to learn the target task from scratch.
Unfortunately, as pointed out by Maurer et al. (2016), there exists an example that satisfies the i.i.d. task assumption for which is unavoidable (or in the realizable setting). This means that the i.i.d. assumption alone is not sufficient if we want to take advantage of a large amount of samples per task. Therefore, a natural question is:
What connections between tasks enable representation learning to utilize all source data?
In this paper, we obtain the first set of results that fully utilize the data from source tasks. We replace the i.i.d. assumption over tasks with natural structural conditions on the input distributions and linear predictors. These conditions depict that the target task can be in some sense “covered” by the source tasks, which will further give rise to the desirable guarantees.
A technical insight coming out of our analysis is that any capacity-controlled method that gets low test error on the source tasks must also get low test error on the target task by virtue of being forced to learn a good representation. Our result on high-dimensional representations and overparametrized neural networks shows that the capacity control for representation learning does not have to be through explicit low dimensionality.
The rest of the paper is organized as follows. We review related work in Section 2. In Section 3, we formally describe the setting we consider. We next present our analysis in four different settings:
Section 4 presents the results for low-dimensional linear representation learning.
Section 5 presents the results for low-dimensional nonlinear representation classes, including neural networks.
Section 6 presents the results for high-dimensional linear representation learning.
Section 7 presents the results for representation learning in overparametrized neural networks.
Finally, we conclude in Section 8 and leave most of the proofs to appendices.
Related Work
The idea of multitask representation learning at least dates back to Caruana (1997), Thrun and Pratt (1998), Baxter (2000). Empirically, representation learning has shown its great power in various domains; see Bengio et al. (2013) for a survey. In particular, representation learning is widely adopted for few-shot learning tasks (Sun et al., 2017, Goyal et al., 2019). Representation learning is also closely connected to meta-learning (Schaul and Schmidhuber, 2010). Recent work Raghu et al. (2019) empirically suggested that the effectiveness of the popular meta-learning algorithm Model Agnostic Meta-Learning (MAML) is due to its ability to learn a useful representation. The scheme we analyze in this paper is closely related to Lee et al. (2019), Bertinetto et al. (2018) for meta-learning.
The concurrent work of Tripuraneni et al. (2020a) studies low-dimensional linear representation learning and obtains a similar result as ours in this case, but they assume isotropic inputs for all tasks, which is a special case of our result. Furthermore, we also provide results for high-dimensional linear representations, general non-linear representations, and overparametrized neural networks. Tripuraneni et al. (2020a) also give a computationally efficient algorithm for standard Gaussian inputs and a lower bound for subspace recovery in the low-dimensional linear setting. Subsequent work of Tripuraneni et al. (2020b) generalizes this work from linear task-specific layers to nonlinear task specific layers and bounded loss functions via proposing a new diversity assumption.
Another recent line of theoretical work analyzed gradient-based meta-learning methods (Denevi et al., 2019, Finn et al., 2019, Khodak et al., 2019) and showed guarantees for convex losses by using tools from online convex optimization. Lastly, we remark that there are analyses for other representation learning schemes (Arora et al., 2019, McNamara and Balcan, 2017, Galanti et al., 2016, Alquier et al., 2016, Denevi et al., 2018).
Notation and Setup
We use the standard , and notation to hide universal constant factors. We also use or to indicate , and use or to mean that for a sufficiently large universal constant .
Let be the returned solution. We are interested in whether our learned predictor works well on average for the target task, i.e., we want the population loss
to be small. In particular, we are interested in the few-shot learning setting, where the number of samples from the target task is small – much smaller than the number of samples required for learning the target task from scratch.
where and are independent. Our goal is to bound the excess risk of our learned model on the target task, i.e., how much our learned model performs worse than the optimal model on the target task:
Low-Dimensional Linear Representations
With this notation, (6) can be rewritten as
With the learned representation from (7), for the target task, we further find a linear function on top of the representation:
There exists such that for all .Note that Assumption 4.2 is a significant generalization of the identically distributed isotropic assumption used in concurrent work Tripuraneni et al. (2020a): they require .
Assumption 4.1 is a standard assumption in statistical learning to obtain probabilistic tail bounds used in our proof. It may be replaced with other moment or boundedness conditions if we adopt different tail bounds in the analysis.
Assumption 4.2 says that every direction spanned by should also be spanned by (), and the parameter quantifies how “easy” it is for to cover . Intuitively, the larger is, the easier it is to cover the target domain using source domains, and we will indeed see that the risk will be proportional to . We remark that we do not necessarily need for all ; as long as this holds for a constant fraction of ’s, our result is valid.
We also make the following assumption that characterizes the diversity of the source tasks.
Finally, we make the following assumption on the distribution of the target task.
Assumption 4.4 can be removed at the cost of a slightly worse risk bound. See Remark 4.2. Our main result in this section is the following theorem.
Fix a failure probability . Under Assumptions 4.1, 4.2, 4.3 and 4.4, we further assume and that the sample sizes in source and target tasks satisfy , , and . Define . Then with probability at least over the samples, the expected excess risk of the learned predictor on the target task satisfies
The proof of Theorem 4.1 is in Appendix A. Theorem 4.1 shows that it is possible to learn the target task using only samples via learning a good representation from the source tasks, which is better than the baseline sample complexity for linear regression, thus demonstrating the benefit of representation learning. It also shows that all samples from source tasks can be pooled together, bypassing the barrier under the i.i.d. tasks assumption.
We note that all our results apply to multi-class problems by removing , with similar class-diversity assumption. Specifically, when source and target have and multi-class labeled samples (instead of independent tasks), using quadratic loss on the one-hot labels, our results apply similarly and will attain an excess risk of the form (see e.g. Lee et al. (2020)). Notice the result is independent of the number of classes.
We can drop Assumption 4.4 and easily obtain the following excess risk bound for any deterministic by slightly modifying the proof of Theorem 4.1:
which is only at most times larger than the bound in (9).
General Low-Dimensional Representations
Now we return to the general case described in Section 3 where we allow a general representation function class . We still assume that the representation is of low dimension . The goal is to obtain a result similar to Theorem 4.1. In this section we assume that inputs from all the tasks follow the same distribution, i.e., , but each task still has its own specialization function (c.f. (4)). We remark that despite this restriction, our result in this section still applies to many interesting and nontrivial scenarios – consider the case where the inputs are all images from ImageNet and each task asks whether the image is from a specific class.
To characterize the complexity of the representation function class , we need the standard definition of Gaussian width.
We will measure the complexity of using the Gaussian width of the following set that depends on the input data :
It is easy to verify for any and .See the proof of Lemma B.1.
We make the following assumptions on the input distribution , which ensure concentration properties of the representation covariances.
where is the empirical distribution over the samples.
where is the empirical distribution over the samples.
Our main theorem in this section is the following:
Theorem 5.1 is very similar to Theorem 4.1 in terms of the result and the assumptions made. In the bound (11), the complexity of is captured by the Gaussian width of the data-dependent set defined in (10). Data-dependent complexity measures are ubiquitous in generalization theory, one of the most notable examples being Rademacher complexity. Similar complexity measure also appeared in existing representation learning theory (Maurer et al., 2016). Usually, for specific examples, we can apply concentration bounds to get rid of the data dependency, such as our result for linear representations (Theorem 4.1).
Our assumptions on the linear specification functions ’s are the same as in Theorem 4.1. The probabilistic assumption on can also be removed at the cost of an additional factor of in the bound – see Remark 4.2.
The proof of Theorem 5.1 is given in Appendix B. Here we prove an important intermediate result on the in-sample risk, which explains how the Gaussian width of arises.
Let and be the optimal solution to (2). Then with probability at least we have
By the optimality of and for (2), we know
Plugging in ( is independent of ), we get
Then the proof is completed using (12). ∎
High-Dimensional Linear Representations
In this section, we consider the case where the representation is a general linear map without an explicit dimensionality constraint, and we will prove a norm-based result by exploiting the intrinsic dimension of the representation. Such a generalization is desirable since in many applications the representation dimension is not restricted.
In this section we additionally assume that all tasks have the same input covariance:
The input distributions in all tasks satisfy .
Note that each task still has its own specialization function (c.f. (4)). We remark that there are many interesting and nontrivial scenarios under Assumption 6.1 – for example, consider the case where the inputs in each task are all images from ImageNet and each task asks whether the image is from a specific class.
Since we do not have a dimensionality constraint, we modify (7) by adding norm constraints:
For the target task, we also modify (8) by adding a norm constraint:
We will specify the choices of regularization, i.e., and in Theorem 6.1.
Fix a failure probability . Under Assumptions 4.1 and 6.1, we further assume , Let , and proper specified in Lemma C.2. Let the target task model be coherent with the source task models in the sense that . Then with probability at least over the samples, the expected excess risk of the learned predictor on the target task satisfies:
The proof of Theorem 6.1 is given in Appendix C. Note that when each is of unit norm. Thus should generally be regarded as for a well-behaved that is nearly low-dimensional. In this regime, Theorem 6.1 indicates that we are able to exploit all samples from the source tasks, similar to Theorem 4.1.
With a good representation, the sample complexity on the target task can also improve over learning the target task from scratch. Consider the baseline of regular ridge regression directly applied to the target task data:
Although the optimization problem (13) is non-convex, its structure allows us to apply existing landscape analysis of matrix factorization problems (Haeffele et al., 2014) and to show that it has the nice properties of no strict saddles and no bad local minima. Therefore, randomly initialized gradient descent or perturbed gradient descent are guaranteed to converge to a global minimum of (13) (Ge et al., 2015, Lee et al., 2016, Jin et al., 2017).
Neural Networks
In this section, we show that we can provably learn good representations in a neural network.
On the target task, we simply re-train the output layer while fixing the hidden layer weights:
Fix a failure probability . Under Assumptions 4.1, 7.1 and 7.2, let , . Let the target task model be coherent with the source task models in the sense that . Set .Then with probability at least over the samples, the expected excess risk of the learned predictor on the target task satisfies:
To highlight the advantage of representation learning, we compare to training a neural network with weight decay directly on the target task:
The error of the baseline method in fixed-design is
We see that Equation (20) is always smaller than Equation (22) since . See Appendix D for the proof of Theorem 7.1 and the calculation of (22).
Conclusion
We gave the first statistical analysis showing that representation learning can fully exploit all data points from source tasks to enable few-shot learning on a target task. This type of results were shown for both low-dimensional and high-dimensional representation function classes.
There are many important directions to pursue in representation learning and few-shot learning. Our results in Sections 6 and 7 indicate that explicit low dimensionality is not necessary, and norm-based capacity control also forces the classifier to learn good representations. Further questions include whether this is a general phenomenon in all deep learning models, whether other capacity control can be applied, and how to optimize to attain good representations.
Acknowledgments
SSD acknowledges support of National Science Foundation (Grant No. DMS-1638352) and the Infosys Membership. JDL acknowledges support of the ARO under MURI Award W911NF-11-1-0303, the Sloan Research Fellowship, and NSF CCF 2002272. WH is supported by NSF, ONR, Simons Foundation, Schmidt Foundation, Amazon Research, DARPA and SRC. QL is supported by NSF #2030859 and the Computing Research Association for the CIFellows Project. The authors also acknowledge the generous support of the Institute for Advanced Study on the Theoretical Machine Learning program, where SSD, WH, JDL, and QL were participants.
References
Appendix A Proof of Theorem 4.1
We first prove several claims and then combine them to finish the proof of Theorem 4.1. We will use technical lemmas proved in Section A.1.
Suppose for . Then with probability at least over the inputs in the source tasks, we have
The proof is finished by taking a union bound over all . ∎
Since and , the above inequality becomes
Under the setting of Theorem 4.1, with probability at least we have
We assume that (23) is true, which happens with probability at least according to Claim A.1.
Let and . From the optimality of and for (7) we have . Plugging in , this becomes
Next we give a high-probability upper bound on using the randomness in . Since ’s depend on which depends on , we will need an -net argument to cover all possible . First, for any fixed , we let where . The ’s defined in this way are independent of . Since has i.i.d. entries, we know that is distributed as . Using the standard tail bound for random variables, we know that with probability at least over ,
Therefore, using the same argument in (A) we know that with probability at least ,
Now, from Lemma A.5 we know that there exists an -net of in Frobenius norm such that and . Applying a union bound over , we know that with probability at least ,
Choosing , we know that (28) holds with probability at least .
We will use (23), (26) and (28) to complete the proof of the claim. This is done in the following steps:
Since , we know that with probability at least ,
Upper bounding .
From (26) we have , which implies . On the other hand, letting the -th column of be , we have
where . Hence we obtain
Applying the -net .
Let such that . Then we have
We have the following chain of inequalities:
Finally, we let , and recall . Then the above inequality implies
The high-probability events we have used in the proof are (23), (28) and (29). By a union bound, the failure probability is at most . Therefore the proof is completed. ∎
Under the setting of Theorem 4.1, with probability at least , we have
From the optimality of and in (6) we know for each . Then we have
Next, we write and . Recall that we have from Claim A.2. Then using Lemma A.7 we can obtain
For the target task, the excess risk of our learned linear predictor is
Applying Claim A.2 with , we have
which implies for . This becomes
From the optimality of in (8) we know . It follows that
For the second term above, notice that , and thus with probability at least we have . Therefore we obtain the final bound
where the last inequality is due to . ∎
Finally, we need to transform into an -net that is a subset of . This can be done by projecting each point in onto . Namely, for each , let be its closest point in (in Frobenium norm); then define . Then we have and is an -net of , because for any , there exists such that , which implies and . ∎
Let . Then it suffices to show with probability at least .
Next, take a -net of with size . By a union bound over all , we have
Plugging in and noticing , the above inequality becomes
where the last inequality is due to .
Therefore, with probability at least we have . Suppose this indeed happens. Next, for any , there exists such that . Then we have
Taking a supreme over , we obtain , i.e., . ∎
If two matrices and (with the same number of columns) satisfy , then for any matrix (of compatible dimensions), we have
As a consequence, for any matrices and (of compatible dimensions), we have
For the first part of the lemma, it suffices to show the following for any vector :
Let . Then we have
For the second part, from we know
Appendix B Proof of Theorem 5.1
The proof is conditioned on several high-probability events, each happening with probability at least . By a union bound at the end, the final success probability is also at least . We can always rescale by a constant factor such that the final probability is at least . Therefore, we will not carefully track the constants before in the proof. All the ’s should be understood as .
We use the following notion of representation divergence.
It is easy to verify , for any and . See Lemma B.1’s proof.
The next lemma shows a relation between (symmetric) covariance and divergence.
Similarly, letting , we have
Under the setting of Theorem 5.1, with probability at least we have
We continue to use the notation from Claim 5.3 and its proof.
Let be the empirical distribution over the samples in (). According to Assumptions 5.1 and 5.2 as well as the setting in Theorem 5.1, we know that the followings are satisfied with probability at least :
By the optimality of and for (2), we know . Then we have the following chain of inequalities:
Now we can finish the proof of Theorem 5.1.
Taking expectation over , we get
Appendix C Proof of Theorem 6.1
Let . Recall and are derived from Eqn. (13) and let . We first note that the constraint set ensures and at global minimum. On the other hand, our constraint for is also expressive enough to attain any that satisfies . See reference e.g. Srebro and Shraibman (2005). Therefore at global minimum and .
For the ease of proof, we introduce the following auxiliary functions and parameters. Write
We define terms and that will be used to bound intrinsic dimension concentration error in the input signal. Namely with high probability, , and similarly . Additionally we use to bound the estimation error (for fixed design) incurred when using noisy label and .
The choice of , and are respectively justified in Lemma C.5, Claim C.4, Lemma C.10 and Claim C.11, along with some more detailed descriptions.
Each step is with high probability over the randomness of or . Therefore overall by union bound, with probability , by plugging in the values of and we have:
Notice a term is absorbed by since we assume . ∎
and , for any .
Here is the adjoint operator of such that .
With the optimality of we have:
Let . Therefore
Therefore , and clearly both terms satisfy and .
For a fixed , let , we have
where .
With the definition of we write the basic inequality:
C.2 Technical Lemmas
This section includes the technical details for several parts: bounding the noise term from basic inequality; and intrinsic dimension concentration for both source and target tasks.
We use matrix Bernstein with intrinsic dimension to bound (See Theorem 7.3.1 in Tropp et al. (2015)).
Write .
Then from intrinsic matrix bernstein (Theorem 7.3.1 in Tropp et al. (2015)), with probability we have, , which gives
The sub-gaussian norm of some vector is defined as:
are called the Gaussian width of and the Gaussian complexity of , respectively.
where is the maximal sub-gaussian norm of the rows of . A high-probability version states as follows. With probability ,
where the radius .
Appendix D Proof of Theorem 7.1
Let be the value of Equation (18) when the network has neurons and be the value of Equation (39). Then
Let be solutions to Equation (18). Let and be a diagonal matrix whose entries are . The network and it satisfies
We first show that . Define . We verify that
Due to the regularizer, and using the AM-GM inequality, at optimality . Next, we verify that the two regularizer values are the same. Let be the -th row vector of . We have
Thus the network given by has the same network outputs and regularizer values. Thus .
Finally, we show that . Let for be the support of the optimal measure of (39). Define , where is a matrix whose rows are , and such that .
Finally by our construction , so the regularizer values agree. Thus .
where and . With these in place, we note that Equation (39) can be expressed as Equation (13) with constrained to be a diagonal operator and as the lifted features .
The global minimizer of Equation (39) with may have infinite support, so the corresponding value may not be achieved by minimizing (18). However, Theorem 6.1 only requires that the we obtain a learner network with regularized loss less than the regularized loss of the teacher network. Since the teacher network has neurons, this value is attainable by (18). Thus the finite-size network does not need to attain the global minimum of (39) for Claim C.1 to apply.
Since Theorem 6.1 has no dependence (even in the logarithmic terms) on the input dimension of the data, it can be applied when the input features the infinite-dimensional feature vector . The only part of the proof of Theorem 6.1 specific to the nuclear norm is that the dual norm is the operator norm. In Lemma C.5 we had an upper bound on . Since we use the norm, we must upper bound , the dual of the -norm. Note that , so the upper bound in Lemma C.5 still applies. Thus, Theorem 7.1 follows from Theorem 6.1. ∎