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 Ω(1T)\Omega(\frac{1}{\sqrt{T}}) is unavoidable (or Ω(1T)\Omega(\frac{1}{T}) 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 n1Tn_{1}T 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 wtw_{t} 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 O(⋅)O(\cdot), Ω(⋅)\Omega(\cdot) and Θ(⋅)\Theta(\cdot) notation to hide universal constant factors. We also use a≲ba\lesssim b or b≳ab\gtrsim a to indicate a=O(b)a=O(b), and use a≫ba\gg b or b≪ab\ll a to mean that a≥C⋅ba\geq C\cdot b for a sufficiently large universal constant C>0C>0.

Let w^T+1\hat{{\bm{w}}}_{T+1} be the returned solution. We are interested in whether our learned predictor x↦⟨w^T+1,ϕ^(x)⟩{\bm{x}}\mapsto\langle\hat{{\bm{w}}}_{T+1},{\hat{\phi}}({\bm{x}})\rangle 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 n2n_{2} from the target task is small – much smaller than the number of samples required for learning the target task from scratch.

where x{\bm{x}} and zz 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 (ϕ^,w^T+1)({\hat{\phi}},\hat{{\bm{w}}}_{T+1}) performs worse than the optimal model (ϕ∗,wT+1∗)(\phi^{*},{\bm{w}}_{T+1}^{*}) on the target task:

Low-Dimensional Linear Representations

With this notation, (6) can be rewritten as

With the learned representation B^\hat{B} from (7), for the target task, we further find a linear function on top of the representation:

There exists c>0c>0 such that Σt⪰c⋅ΣT+1\Sigma_{t}\succeq c\cdot\Sigma_{T+1} for all t∈[T]t\in[T].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 Σ1=Σ2=⋯=ΣT+1=I\Sigma_{1}=\Sigma_{2}=\cdots=\Sigma_{T+1}=I.

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 ΣT+1\Sigma_{T+1} should also be spanned by Σt\Sigma_{t} (t∈[T]t\in[T]), and the parameter cc quantifies how “easy” it is for Σt\Sigma_{t} to cover ΣT+1\Sigma_{T+1}. Intuitively, the larger cc 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 1c\frac{1}{c}. We remark that we do not necessarily need Σt⪰c⋅ΣT+1\Sigma_{t}\succeq c\cdot\Sigma_{T+1} for all t∈[T]t\in[T]; as long as this holds for a constant fraction of tt’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 δ∈(0,1)\delta\in(0,1). Under Assumptions 4.1, 4.2, 4.3 and 4.4, we further assume 2k≤min⁡{d,T}2k\leq\min\{d,T\} and that the sample sizes in source and target tasks satisfy n1≫ρ4(d+log⁡Tδ)n_{1}\gg\rho^{4}(d+\log\frac{T}{\delta}), n2≫ρ4(k+log⁡1δ)n_{2}\gg\rho^{4}(k+\log\frac{1}{\delta}), and cn1≥n2cn_{1}\geq n_{2}. Define κ=max⁡t∈[T]λmax⁡(Σt)min⁡t∈[T]λmin⁡(Σt)\kappa=\frac{\max_{t\in[T]}\lambda_{\max}(\Sigma_{t})}{\min_{t\in[T]}\lambda_{\min}(\Sigma_{t})}. Then with probability at least 1−δ1-\delta over the samples, the expected excess risk of the learned predictor x↦w^T+1⊤B^x{\bm{x}}\mapsto\hat{{\bm{w}}}_{T+1}^{\top}\hat{B}{\bm{x}} 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 O(k)O(k) samples via learning a good representation from the source tasks, which is better than the baseline O(d)O(d) sample complexity for linear regression, thus demonstrating the benefit of representation learning. It also shows that all n1Tn_{1}T samples from source tasks can be pooled together, bypassing the Ω(1T)\Omega(\frac{1}{T}) barrier under the i.i.d. tasks assumption.

We note that all our results apply to multi-class problems by removing TT, with similar class-diversity assumption. Specifically, when source and target have n1n_{1} and n2n_{2} 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 σ2O(kdlog⁡(κn1)cn1+k+log⁡1/δn2)\sigma^{2}O\left(\frac{kd\log(\kappa n_{1})}{cn_{1}}+\frac{k+\log 1/\delta}{n_{2}}\right) (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 wT+1∗{\bm{w}}_{T+1}^{*} by slightly modifying the proof of Theorem 4.1:

which is only at most kk 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 Φ\Phi. We still assume that the representation is of low dimension kk. 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., p1=⋯=pT+1=pp_{1}=\cdots=p_{T+1}=p, but each task tt still has its own specialization function wt∗{\bm{w}}_{t}^{*} (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 Φ\Phi, we need the standard definition of Gaussian width.

We will measure the complexity of Φ\Phi using the Gaussian width of the following set that depends on the input data X\mathcal{X}:

It is easy to verify Λq(ϕ,ϕ′)⪰0\Lambda_{q}(\phi,\phi^{\prime})\succeq 0 for any ϕ,ϕ′\phi,\phi^{\prime} and qq.See the proof of Lemma B.1.

We make the following assumptions on the input distribution pp, which ensure concentration properties of the representation covariances.

where p^\hat{p} is the empirical distribution over the nn samples.

where p^\hat{p} is the empirical distribution over the nn 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 Φ\Phi is captured by the Gaussian width of the data-dependent set FX(Φ)\mathcal{F}_{\mathcal{X}}(\Phi) 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 wt∗{\bm{w}}_{t}^{*}’s are the same as in Theorem 4.1. The probabilistic assumption on wT+1∗{\bm{w}}_{T+1}^{*} can also be removed at the cost of an additional factor of kk 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 FX(Φ)\mathcal{F}_{\mathcal{X}}(\Phi) arises.

Let ϕ^\hat{\phi} and w^1,…,w^T\hat{{\bm{w}}}_{1},\ldots,\hat{{\bm{w}}}_{T} be the optimal solution to (2). Then with probability at least 1−δ1-\delta we have

By the optimality of ϕ^\hat{\phi} and w^1,…,w^T\hat{{\bm{w}}}_{1},\ldots,\hat{{\bm{w}}}_{T} for (2), we know

Plugging in yt=ϕ∗(Xt)wt∗+zt{\bm{y}}_{t}=\phi^{*}(X_{t}){\bm{w}}^{*}_{t}+{\bm{z}}_{t} (zt∼N(0,I){\bm{z}}_{t}\sim\mathcal{N}(0,I) is independent of XtX_{t}), 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 Σ1=⋯=ΣT+1=Σ\Sigma_{1}=\cdots=\Sigma_{T+1}=\Sigma.

Note that each task tt still has its own specialization function wt∗{\bm{w}}_{t}^{*} (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., λ\lambda and rr in Theorem 6.1.

Fix a failure probability δ∈(0,1)\delta\in(0,1). Under Assumptions 4.1 and 6.1, we further assume n1≥n2n_{1}\geq n_{2}, R=∥Θ∥∗.R=\|\Theta\|_{*}. Let r=2R/Tr=2\sqrt{R/T}, Rˉ=R/T\bar{R}=R/\sqrt{T} and proper λ\lambda specified in Lemma C.2. Let the target task model θT+1∗{\bm{\theta}}_{T+1}^{*} be coherent with the source task models Θ∗\Theta^{*} in the sense that θT+1∗∼ν=N(0,Θ∗(Θ∗)⊤/T){\bm{\theta}}_{T+1}^{*}\sim\nu=\mathcal{N}(\bm{0},\Theta^{*}(\Theta^{*})^{\top}/T). Then with probability at least 1−δ1-\delta over the samples, the expected excess risk of the learned predictor x↦w^T+1⊤B^⊤x{\bm{x}}\mapsto\hat{{\bm{w}}}_{T+1}^{\top}\hat{B}^{\top}{\bm{x}} on the target task satisfies:

The proof of Theorem 6.1 is given in Appendix C. Note that ∥Θ∗∥F=T\|\Theta^{*}\|_{F}=\sqrt{T} when each θt∗{\bm{\theta}}^{\ast}_{t} is of unit norm. Thus Rˉ=∥Θ∗∥∗/T\bar{R}=\left\|\Theta^{*}\right\|_{*}/\sqrt{T} should generally be regarded as O(1)O(1) for a well-behaved Θ∗\Theta^{*} that is nearly low-dimensional. In this regime, Theorem 6.1 indicates that we are able to exploit all n1Tn_{1}T 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 δ∈(0,1)\delta\in(0,1). Under Assumptions 4.1, 7.1 and 7.2, let n1≥n2n_{1}\geq n_{2}, Rˉ=(12∥B∗∥F2+12∥W∗∥F2)/T\bar{R}=(\frac{1}{2}\|B^{*}\|_{F}^{2}+\frac{1}{2}\|W^{*}\|_{F}^{2})/\sqrt{T}. Let the target task model fαT+1=⟨αT+1,ϕ(x)⟩f_{\alpha_{T+1}}=\langle\alpha_{T+1},\phi({\bm{x}})\rangle be coherent with the source task models in the sense that αT+1∗∼ν\alpha_{T+1}^{*}\sim\nu. Set r2=(∥B∗∥F2+∥W∗∥F2)/Tr^{2}=(\|B^{*}\|_{F}^{2}+\|W^{*}\|_{F}^{2})/T.Then with probability at least 1−δ1-\delta over the samples, the expected excess risk of the learned predictor x↦w^T+1⊤(B^⊤x)+{\bm{x}}\mapsto\hat{{\bm{w}}}_{T+1}^{\top}(\hat{B}^{\top}{\bm{x}})_{+} 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 n1T≥n2n_{1}T\geq n_{2}. 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 n1≫ρ4(d+log⁡(T/δ))n_{1}\gg\rho^{4}(d+\log(T/\delta)) for δ∈(0,1)\delta\in(0,1). Then with probability at least 1−δ101-\frac{\delta}{10} over the inputs X1,…,XTX_{1},\ldots,X_{T} in the source tasks, we have

The proof is finished by taking a union bound over all t∈[T]t\in[T]. ∎

Since 1n2VDU⊤XˉT+1⊤XˉT+1UDV⊤=1n2B⊤ΣT+11/2XˉT+1⊤XˉT+1ΣT+11/2B=1n2B⊤XT+1⊤XT+1B\frac{1}{n_{2}}VDU^{\top}\bar{X}_{T+1}^{\top}\bar{X}_{T+1}UDV^{\top}=\frac{1}{n_{2}}B^{\top}\Sigma_{T+1}^{1/2}\bar{X}_{T+1}^{\top}\bar{X}_{T+1}\Sigma_{T+1}^{1/2}B=\frac{1}{n_{2}}B^{\top}X_{T+1}^{\top}X_{T+1}B and VDDV⊤=VDU⊤UDV⊤=B⊤ΣT+1BVDDV^{\top}=VDU^{\top}UDV^{\top}=B^{\top}\Sigma_{T+1}B, the above inequality becomes

Under the setting of Theorem 4.1, with probability at least 1−δ51-\frac{\delta}{5} we have

We assume that (23) is true, which happens with probability at least 1−δ101-\frac{\delta}{10} according to Claim A.1.

Let Θ^=B^W^\hat{\Theta}=\hat{B}\hat{W} and Θ∗=B∗W∗\Theta^{*}=B^{*}W^{*}. From the optimality of B^\hat{B} and W^\hat{W} for (7) we have ∥Y−X(Θ^)∥F2≤∥Y−X(Θ∗)∥F2\|Y-\mathcal{X}(\hat{\Theta})\|_{F}^{2}\leq\|Y-\mathcal{X}(\Theta^{*})\|_{F}^{2}. Plugging in Y=X(Θ∗)+ZY=\mathcal{X}(\Theta^{*})+Z, this becomes

Next we give a high-probability upper bound on ∑t=1T∥Ut⊤zt∥2\sum_{t=1}^{T}\left\|U_{t}^{\top}{\bm{z}}_{t}\right\|^{2} using the randomness in ZZ. Since UtU_{t}’s depend on VV which depends on ZZ, we will need an ϵ\epsilon-net argument to cover all possible V∈Od,2kV\in\mathcal{O}_{d,2k}. First, for any fixed Vˉ∈Od,2k\bar{V}\in\mathcal{O}_{d,2k}, we let XtVˉ=UˉtQˉtX_{t}\bar{V}=\bar{U}_{t}\bar{Q}_{t} where Uˉt∈On,2k\bar{U}_{t}\in\mathcal{O}_{n,2k}. The Uˉt\bar{U}_{t}’s defined in this way are independent of ZZ. Since ZZ has i.i.d. N(0,σ2)\mathcal{N}(0,\sigma^{2}) entries, we know that σ−2∑t=1T∥Uˉt⊤zt∥2\sigma^{-2}\sum_{t=1}^{T}\left\|\bar{U}_{t}^{\top}{\bm{z}}_{t}\right\|^{2} is distributed as χ2(2kT)\chi^{2}(2kT). Using the standard tail bound for χ2\chi^{2} random variables, we know that with probability at least 1−δ′1-\delta^{\prime} over ZZ,

Therefore, using the same argument in (A) we know that with probability at least 1−δ′1-\delta^{\prime},

Now, from Lemma A.5 we know that there exists an ϵ\epsilon-net N\mathcal{N} of Od,2k\mathcal{O}_{d,2k} in Frobenius norm such that N⊂Od,2k\mathcal{N}\subset\mathcal{O}_{d,2k} and ∣N∣≤(62kϵ)2kd|\mathcal{N}|\leq(\frac{6\sqrt{2k}}{\epsilon})^{2kd}. Applying a union bound over N\mathcal{N}, we know that with probability at least 1−δ′∣N∣1-\delta^{\prime}|\mathcal{N}|,

Choosing δ′=δ20(62kϵ)2kd\delta^{\prime}=\frac{\delta}{20(\frac{6\sqrt{2k}}{\epsilon})^{2kd}}, we know that (28) holds with probability at least 1−δ201-\frac{\delta}{20}.

We will use (23), (26) and (28) to complete the proof of the claim. This is done in the following steps:

Since σ−2∥Z∥F2∼χ2(n1T)\sigma^{-2}\left\|Z\right\|_{F}^{2}\sim\chi^{2}(n_{1}T), we know that with probability at least 1−δ201-\frac{\delta}{20},

Upper bounding ∥Δ∥F\left\|\Delta\right\|_{F}.

From (26) we have ∥X(Δ)∥F2≤2∥Z∥F∥X(Δ)∥F\left\|\mathcal{X}(\Delta)\right\|_{F}^{2}\leq 2\left\|Z\right\|_{F}\left\|\mathcal{X}(\Delta)\right\|_{F}, which implies ∥X(Δ)∥F≤2∥Z∥F≲σn1T+log⁡(1/δ)\left\|\mathcal{X}(\Delta)\right\|_{F}\leq 2\left\|Z\right\|_{F}\lesssim\sigma\sqrt{n_{1}T+\log(1/\delta)}. On the other hand, letting the tt-th column of Δ\Delta be δt{\bm{\delta}}_{t}, we have

where λ‾=min⁡t∈[T]λmin⁡(Σt){\underline{\lambda}}=\min_{t\in[T]}\lambda_{\min}(\Sigma_{t}). Hence we obtain

Applying the ϵ\epsilon-net N\mathcal{N}.

Let Vˉ∈N\bar{V}\in\mathcal{N} such that ∥V−Vˉ∥F≤ϵ\left\|V-\bar{V}\right\|_{F}\leq\epsilon. Then we have

We have the following chain of inequalities:

Finally, we let ϵ=kκn1\epsilon=\frac{k}{\sqrt{\kappa}n_{1}}, and recall δ′=δ20(62kϵ)2kd\delta^{\prime}=\frac{\delta}{20(\frac{6\sqrt{2k}}{\epsilon})^{2kd}}. 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 δ10+δ20+δ20=δ5\frac{\delta}{10}+\frac{\delta}{20}+\frac{\delta}{20}=\frac{\delta}{5}. Therefore the proof is completed. ∎

Under the setting of Theorem 4.1, with probability at least 1−2δ51-\frac{2\delta}{5}, we have

From the optimality of B^\hat{B} and W^\hat{W} in (6) we know XtB^w^t=PXtB^yt=PXtB^(XtB∗wt∗+zt)X_{t}\hat{B}\hat{{\bm{w}}}_{t}=P_{X_{t}\hat{B}}{\bm{y}}_{t}=P_{X_{t}\hat{B}}(X_{t}B^{*}{\bm{w}}_{t}^{*}+{\bm{z}}_{t}) for each t∈[T]t\in[T]. Then we have

Next, we write B^=[B^,B∗][I0]=:BA\hat{B}=[\hat{B},B^{*}]\begin{bmatrix}I\\ 0\end{bmatrix}=:BA and B∗=[B^,B∗][0I]=:BCB^{*}=[\hat{B},B^{*}]\begin{bmatrix}0\\ I\end{bmatrix}=:BC. Recall that we have 1n2B⊤XT+1⊤XT+1B⪯1.1B⊤ΣT+1B\frac{1}{n_{2}}B^{\top}X_{T+1}^{\top}X_{T+1}B\preceq 1.1B^{\top}\Sigma_{T+1}B from Claim A.2. Then using Lemma A.7 we can obtain

For the target task, the excess risk of our learned linear predictor x↦(B^w^T+1)⊤x{\bm{x}}\mapsto(\hat{B}\hat{{\bm{w}}}_{T+1})^{\top}{\bm{x}} is

Applying Claim A.2 with B=[B^,B∗]B=[\hat{B},B^{*}], we have

which implies 0.9v⊤B⊤ΣT+1Bv≤1n2vB⊤XT+1⊤XT+1Bv0.9{\bm{v}}^{\top}B^{\top}\Sigma_{T+1}B{\bm{v}}\leq\frac{1}{n_{2}}{\bm{v}}B^{\top}X_{T+1}^{\top}X_{T+1}B{\bm{v}} for v=[w^T+1wT+1∗]{\bm{v}}=\begin{bmatrix}\hat{{\bm{w}}}_{T+1}\\ {\bm{w}}_{T+1}^{*}\end{bmatrix}. This becomes

From the optimality of w^T+1\hat{{\bm{w}}}_{T+1} in (8) we know XT+1B^w^T+1=PXT+1B^yT+1=PXT+1B^(XT+1B∗wT+1∗+zT+1)X_{T+1}\hat{B}\hat{{\bm{w}}}_{T+1}=P_{X_{T+1}\hat{B}}{\bm{y}}_{T+1}=P_{X_{T+1}\hat{B}}(X_{T+1}B^{*}{\bm{w}}_{T+1}^{*}+{\bm{z}}_{T+1}). It follows that

For the second term above, notice that 1σ2∥PXT+1B^zT+1∥F2∼χ2(k)\frac{1}{\sigma^{2}}\left\|P_{X_{T+1}\hat{B}}{\bm{z}}_{T+1}\right\|_{F}^{2}\sim\chi^{2}(k), and thus with probability at least 1−δ51-\frac{\delta}{5} we have 1σ2∥PXT+1B^zT+1∥F2≲k+log⁡1δ\frac{1}{\sigma^{2}}\left\|P_{X_{T+1}\hat{B}}{\bm{z}}_{T+1}\right\|_{F}^{2}\lesssim k+\log\frac{1}{\delta}. Therefore we obtain the final bound

where the last inequality is due to cn1≥n2cn_{1}\geq n_{2}. ∎

Finally, we need to transform N′\mathcal{N}^{\prime} into an ϵ\epsilon-net N\mathcal{N} that is a subset of Od1,d2\mathcal{O}_{d_{1},d_{2}}. This can be done by projecting each point in N′\mathcal{N}^{\prime} onto Od1,d2\mathcal{O}_{d_{1},d_{2}}. Namely, for each Vˉ∈N′\bar{V}\in\mathcal{N}^{\prime}, let P(Vˉ)\mathcal{P}(\bar{V}) be its closest point in Od1,d2\mathcal{O}_{d_{1},d_{2}} (in Frobenium norm); then define N={P(Vˉ)∣Vˉ∈N′}\mathcal{N}=\{\mathcal{P}(\bar{V})\mid\bar{V}\in\mathcal{N}^{\prime}\}. Then we have ∣N∣≤∣N′∣≤(6d2ϵ)d1d2|\mathcal{N}|\leq|\mathcal{N}^{\prime}|\leq(\frac{6\sqrt{d_{2}}}{\epsilon})^{d_{1}d_{2}} and N\mathcal{N} is an ϵ\epsilon-net of Od1,d2\mathcal{O}_{d_{1},d_{2}}, because for any V∈Od1,d2V\in\mathcal{O}_{d_{1},d_{2}}, there exists Vˉ∈N′\bar{V}\in\mathcal{N}^{\prime} such that ∥V−Vˉ∥F≤ϵ2\left\|V-\bar{V}\right\|_{F}\leq\frac{\epsilon}{2}, which implies P(Vˉ)∈N\mathcal{P}(\bar{V})\in\mathcal{N} and ∥V−P(Vˉ)∥F≤∥V−Vˉ∥F+∥Vˉ−P(Vˉ)∥F≤∥V−Vˉ∥F+∥Vˉ−V∥F=2∥V−Vˉ∥F≤ϵ\left\|V-\mathcal{P}(\bar{V})\right\|_{F}\leq\left\|V-\bar{V}\right\|_{F}+\left\|\bar{V}-\mathcal{P}(\bar{V})\right\|_{F}\leq\left\|V-\bar{V}\right\|_{F}+\left\|\bar{V}-V\right\|_{F}=2\left\|V-\bar{V}\right\|_{F}\leq\epsilon. ∎

Let A=1n∑i=1naiai⊤−IA=\frac{1}{n}\sum_{i=1}^{n}{\bm{a}}_{i}{\bm{a}}_{i}^{\top}-I. Then it suffices to show ∥A∥≤0.1\left\|A\right\|\leq 0.1 with probability at least 1−δ1-\delta.

Next, take a 15\frac{1}{5}-net N⊂Sd−1\mathcal{N}\subset\mathcal{S}^{d-1} of Sd−1\mathcal{S}^{d-1} with size ∣N∣≤eO(d)|\mathcal{N}|\leq e^{O(d)}. By a union bound over all v∈N{\bm{v}}\in\mathcal{N}, we have

Plugging in ϵ=120\epsilon=\frac{1}{20} and noticing ρ>1\rho>1, the above inequality becomes

where the last inequality is due to n≫ρ4(d+log⁡(1/δ))n\gg\rho^{4}\left(d+\log(1/\delta)\right).

Therefore, with probability at least 1−δ1-\delta we have max⁡v∈N∣v⊤Av∣≤120\max_{{\bm{v}}\in\mathcal{N}}|{\bm{v}}^{\top}A{\bm{v}}|\leq\frac{1}{20}. Suppose this indeed happens. Next, for any u∈Sd−1{\bm{u}}\in\mathcal{S}^{d-1}, there exists u′∈N{\bm{u}}^{\prime}\in\mathcal{N} such that ∥u−u′∥≤15\left\|{\bm{u}}-{\bm{u}}^{\prime}\right\|\leq\frac{1}{5}. Then we have

Taking a supreme over u∈Sd−1u\in\mathcal{S}^{d-1}, we obtain ∥A∥≤120+12∥A∥\left\|A\right\|\leq\frac{1}{20}+\frac{1}{2}\left\|A\right\|, i.e., ∥A∥≤110\left\|A\right\|\leq\frac{1}{10}. ∎

If two matrices A1A_{1} and A2A_{2} (with the same number of columns) satisfy A1⊤A1⪰A2⊤A2A_{1}^{\top}A_{1}\succeq A_{2}^{\top}A_{2}, then for any matrix BB (of compatible dimensions), we have

As a consequence, for any matrices BB and B′B^{\prime} (of compatible dimensions), we have

For the first part of the lemma, it suffices to show the following for any vector v{\bm{v}}:

Let w∗∈arg min⁡w∥A1Bw−A1v∥22{\bm{w}}^{*}\in\operatorname*{arg\,min}_{{\bm{w}}}\|A_{1}B{\bm{w}}-A_{1}{\bm{v}}\|_{2}^{2}. Then we have

For the second part, from A1⊤PA1B⊥A1⪰A2⊤PA2B⊥A2A_{1}^{\top}P^{\perp}_{A_{1}B}A_{1}\succeq A_{2}^{\top}P^{\perp}_{A_{2}B}A_{2} we know

Appendix B Proof of Theorem 5.1

The proof is conditioned on several high-probability events, each happening with probability at least 1−Ω(δ)1-\Omega(\delta). By a union bound at the end, the final success probability is also at least 1−Ω(δ)1-\Omega(\delta). We can always rescale δ\delta by a constant factor such that the final probability is at least 1−δ1-\delta. Therefore, we will not carefully track the constants before δ\delta in the proof. All the δ\delta’s should be understood as Ω(δ)\Omega(\delta).

We use the following notion of representation divergence.

It is easy to verify Dq(ϕ,ϕ′)⪰0D_{q}(\phi,\phi^{\prime})\succeq 0, Dq(ϕ,ϕ)=0D_{q}(\phi,\phi)=0 for any ϕ,ϕ′\phi,\phi^{\prime} and qq. See Lemma B.1’s proof.

The next lemma shows a relation between (symmetric) covariance and divergence.

Similarly, letting g(w)=[w⊤,−v⊤]Λq′(ϕ,ϕ′)[w−v]g({\bm{w}})=[{\bm{w}}^{\top},-{\bm{v}}^{\top}]\Lambda_{q^{\prime}}(\phi,\phi^{\prime})\begin{bmatrix}{\bm{w}}\\ -{\bm{v}}\end{bmatrix}, we have

Under the setting of Theorem 5.1, with probability at least 1−δ1-\delta we have

We continue to use the notation from Claim 5.3 and its proof.

Let p^t\hat{p}_{t} be the empirical distribution over the samples in XtX_{t} (t∈[T+1]t\in[T+1]). 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 1−δ1-\delta:

By the optimality of ϕ^\hat{\phi} and w^1,…,w^T\hat{{\bm{w}}}_{1},\ldots,\hat{{\bm{w}}}_{T} for (2), we know ϕ^(Xt)w^t=Pϕ^(Xt)yt=Pϕ^(Xt)(ϕ∗(Xt)wt∗+zt)\hat{\phi}(X_{t})\hat{{\bm{w}}}_{t}=P_{\hat{\phi}(X_{t})}{\bm{y}}_{t}=P_{\hat{\phi}(X_{t})}(\phi^{*}(X_{t}){\bm{w}}^{*}_{t}+{\bm{z}}_{t}). Then we have the following chain of inequalities:

Now we can finish the proof of Theorem 5.1.

Taking expectation over wT+1∗∼ν{\bm{w}}_{T+1}^{*}\sim\nu, we get

Appendix C Proof of Theorem 6.1

Let R=∥Θ∗∥∗R=\|\Theta^{*}\|_{*}. Recall B^\hat{B} and W^\hat{W} are derived from Eqn. (13) and let Θ^:=B^W^\hat{\Theta}:=\hat{B}\hat{W}. We first note that the constraint set {∥w∥i2≤R/T,∥B∥F2≤R}\{\|{\bm{w}}\|_{i}^{2}\leq R/T,\|B\|_{F}^{2}\leq R\} ensures ∥W∥F2≤R\|W\|_{F}^{2}\leq R and ∥WB∥∗≤R\|WB\|_{*}\leq R at global minimum. On the other hand, our constraint for W,BW,B is also expressive enough to attain any Θ^\hat{\Theta} that satisfies ∥Θ^∥∗≤R\|\hat{\Theta}\|_{*}\leq R. See reference e.g. Srebro and Shraibman (2005). Therefore at global minimum ∥W^∥F≤R,∥B^∥F≤R\|\hat{W}\|_{F}\leq\sqrt{R},\|\hat{B}\|_{F}\leq\sqrt{R} and ∥Θ^∥∗≤R\|\hat{\Theta}\|_{*}\leq R.

For the ease of proof, we introduce the following auxiliary functions and parameters. Write

We define terms ϵic,1\epsilon_{ic,1} and ϵic,2\epsilon_{ic,2} that will be used to bound intrinsic dimension concentration error in the input signal. Namely with high probability, ∥Σ1/2Θ∥−1/n1∑t=1T∥Xtθt∥2≤ϵic,1∥Θ∥∗\|\Sigma^{1/2}\Theta\|-\sqrt{1/n_{1}\sum_{t=1}^{T}\|X_{t}\theta_{t}\|^{2}}\leq\epsilon_{ic,1}\|\Theta\|_{*}, and similarly ∥Σ1/2B^v∥−1n2∥XB^v∥2≤ϵic,2∥v∥2\|\Sigma^{1/2}\hat{B}{\bm{v}}\|-\sqrt{\frac{1}{n_{2}}\|X\hat{B}{\bm{v}}\|^{2}}\leq\epsilon_{ic,2}\|v\|_{2}. Additionally we use ϵee,i,i∈{1,2}\epsilon_{ee,i},i\in\{1,2\} to bound the estimation error (for fixed design) incurred when using noisy label yT+1{\bm{y}}_{T+1} and YY.

The choice of ϵee,i\epsilon_{ee,i}, and ϵic,i\epsilon_{ic,i} 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 1−δ/101-\delta/10 over the randomness of X\mathcal{X} or XT+1X_{T+1}. Therefore overall by union bound, with probability 1−δ1-\delta, by plugging in the values of ϵic,i\epsilon_{ic,i} and ϵee,i\epsilon_{ee,i} we have:

Notice a term ∥Σ∥/n1\|\Sigma\|/n_{1} is absorbed by ∥Σ∥/n2\|\Sigma\|/n_{2} since we assume n1≥n2n_{1}\geq n_{2}. ∎

and ∥B^∥F2≤3R\|\hat{B}\|_{F}^{2}\leq 3R, ∥W^∥F2≤3R\|\hat{W}\|_{F}^{2}\leq 3R for any λ≥2n∥X∗(Z)∥2\lambda\geq\frac{2}{n}\|\mathcal{X}^{*}(Z)\|_{2}.

Here X∗\mathcal{X}^{*} is the adjoint operator of X\mathcal{X} such that X∗(Z)=∑i=1TXt⊤ztet⊤\mathcal{X}^{*}(Z)=\sum_{i=1}^{T}X_{t}^{\top}{\bm{z}}_{t}{\bm{e}}_{t}^{\top}.

With the optimality of Θ^\hat{\Theta} we have:

Let Δ=Θ^−Θ∗\Delta=\hat{\Theta}-\Theta^{*}. Therefore

Therefore 12n1∥X(Δ)∥F2+λ2∥Θ^∥∗≤32λ∥Θ∗∥∗\frac{1}{2n_{1}}\|\mathcal{X}(\Delta)\|_{F}^{2}+\frac{\lambda}{2}\|\hat{\Theta}\|_{*}\leq\frac{3}{2}\lambda\|\Theta^{*}\|_{*}, and clearly both terms satisfy 1n1∥X(Δ)∥F2≤3λ∥Θ∗∥∗\frac{1}{n_{1}}\|\mathcal{X}(\Delta)\|_{F}^{2}\leq 3\lambda\|\Theta^{*}\|_{*} and ∥Θ^∥∗≤3∥Θ∗∥∗\|\hat{\Theta}\|_{*}\leq 3\|\Theta^{*}\|_{*}.

For a fixed δ>0\delta>0, let λ=ϵee,12+ϵic,12R\lambda=\epsilon_{ee,1}^{2}+\epsilon_{ic,1}^{2}R, we have

where Sλ:=(B^⊤ΣB^+λI)−1B^⊤ΣS_{\lambda}:=(\hat{B}^{\top}\Sigma\hat{B}+\lambda I)^{-1}\hat{B}^{\top}\Sigma.

With the definition of w^\hat{{\bm{w}}} 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 λ\lambda (See Theorem 7.3.1 in Tropp et al. (2015)).

Write A=1nX⊤Z=1n∑t=1TX⊤ztet⊤=:∑t=1TStA=\frac{1}{\sqrt{n}}X^{\top}Z=\frac{1}{\sqrt{n}}\sum_{t=1}^{T}X^{\top}{\bm{z}}_{t}{\bm{e}}_{t}^{\top}=:\sum_{t=1}^{T}S_{t}.

Then from intrinsic matrix bernstein (Theorem 7.3.1 in Tropp et al. (2015)), with probability 1−δ1-\delta we have, ∥A∥≤O(σlog⁡1δvlog⁡(dΣ)+σlog⁡1δLlog⁡(dΣ))\|A\|\leq{\cal{O}}(\sigma\sqrt{\log\frac{1}{\delta}v\log(d_{\Sigma})}+\sigma\log\frac{1}{\delta}L\log(d_{\Sigma})), which gives

The sub-gaussian norm of some vector y{\bm{y}} is defined as:

are called the Gaussian width of TT and the Gaussian complexity of TT, respectively.

where K=max⁡i∥Ai∥ψ2K=\max_{i}\|A_{i}\|_{\psi_{2}} is the maximal sub-gaussian norm of the rows of AA. A high-probability version states as follows. With probability 1−δ1-\delta,

where the radius r(T):=sup⁡x∈T∥x∥2r(T):=\sup_{{\bm{x}}\in T}\|{\bm{x}}\|_{2}.

Appendix D Proof of Theorem 7.1

Let γd\gamma_{d} be the value of Equation (18) when the network has dd neurons and γd⋆\gamma^{\star}_{d} be the value of Equation (39). Then

Let B,WB,W be solutions to Equation (18). Let Bˉ=BDβ−1\bar{B}=BD_{\beta}^{-1} and DβD_{\beta} be a diagonal matrix whose entries are βj=∥B⊤ej∥2\beta_{j}=\|B^{\top}{\bm{e}}_{j}\|_{2}. The network fB,W(x)=W⊤Dβ(Bˉ⊤x)+f_{B,W}({\bm{x}})=W^{\top}D_{\beta}(\bar{B}^{\top}{\bm{x}})_{+} and it satisfies

We first show that γd⋆≤γd\gamma^{\star}_{d}\leq\gamma_{d}. Define αt(bj∥bj∥)=Wtjβj\alpha_{t}(\frac{\bm{b}_{j}}{\|\bm{b}_{j}\|})=W_{tj}\beta_{j}. We verify that

Due to the regularizer, and using the AM-GM inequality, at optimality βj=∥Wej∥2\beta_{j}=\|W{\bm{e}}_{j}\|_{2}. Next, we verify that the two regularizer values are the same. Let wˉj\bar{\bm{w}}_{j} be the jj-th row vector of WW. We have

Thus the network given by αt⊤ϕ(x)\alpha_{t}^{\top}\phi({\bm{x}}) has the same network outputs and regularizer values. Thus γ⋆≤γd\gamma^{\star}\leq\gamma_{d}.

Finally, we show that γd≤γd⋆\gamma_{d}\leq\gamma^{\star}_{d}. Let bˉj\bar{\bm{b}}_{j} for j∈[d]j\in[d] be the support of the optimal measure of (39). Define βj=∥α(bˉj)∥2\beta_{j}=\sqrt{\|\bm{\alpha}(\bar{\bm{b}}_{j})\|_{2}}, B=BˉDβB=\bar{B}D_{\beta} where Bˉ\bar{B} is a matrix whose rows are bˉj\bar{\bm{b}}_{j}, and WW such that Wjt=αt(bˉj)/∥α(bˉj)∥W_{jt}=\alpha_{t}(\bar{\bm{b}}_{j})/\sqrt{\|\bm{\alpha}(\bar{\bm{b}}_{j})\|}.

Finally by our construction βj=∥Wej∥\beta_{j}=\|W{\bm{e}}_{j}\|, so the regularizer values agree. Thus γd=γd⋆\gamma_{d}=\gamma^{\star}_{d}.

where ∥β∥22=∫β(bˉ)2d(bˉ)\|\beta\|_{2}^{2}=\int\beta(\bar{\bm{b}})^{2}d(\bar{\bm{b}}) and ∥W∥F2=∑t∫wt(bˉ)2d(bˉ)\|W\|_{F}^{2}=\sum_{t}\int{\bm{w}}_{t}(\bar{\bm{b}})^{2}d(\bar{\bm{b}}). With these in place, we note that Equation (39) can be expressed as Equation (13) with BB constrained to be a diagonal operator and xitx_{it} as the lifted features ϕ(xit)\phi(x_{it}).

The global minimizer of Equation (39) with d=∞d=\infty 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 dd 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 ϕ(x)\phi({\bm{x}}). 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 1n∥X⊤Z∥2\frac{1}{n}\|X^{\top}Z\|_{2}. Since we use the ∥⋅∥2,1\|\cdot\|_{2,1} norm, we must upper bound 1n∥X⊤Z∥2,∞\frac{1}{n}\|X^{\top}Z\|_{2,\infty} , the dual of the (2,1)(2,1)-norm. Note that ∥A∥2,∞≤∥A∥2\|A\|_{2,\infty}\leq\|A\|_{2}, so the upper bound in Lemma C.5 still applies. Thus, Theorem 7.1 follows from Theorem 6.1. ∎