Meta-Learning and Universality: Deep Representations and Gradient Descent can Approximate any Learning Algorithm

Chelsea Finn, Sergey Levine

Introduction

Deep neural networks that optimize for effective representations have enjoyed tremendous success over human-engineered representations. Meta-learning takes this one step further by optimizing for a learning algorithm that can effectively acquire representations. A common approach to meta-learning is to train a recurrent or memory-augmented model such as a recurrent neural network to take a training dataset as input and then output the parameters of a learner model (Schmidhuber, 1987; Bengio et al., 1992; Li & Malik, 2017a; Andrychowicz et al., 2016). Alternatively, some approaches pass the dataset and test input into the model, which then outputs a corresponding prediction for the test example (Santoro et al., 2016; Duan et al., 2016; Wang et al., 2016; Mishra et al., 2018). Such recurrent models are universal learning procedure approximators, in that they have the capacity to approximately represent any mapping from dataset and test datapoint to label. However, depending on the form of the model, it may lack statistical efficiency.

In contrast to the aforementioned approaches, more recent work has proposed methods that include the structure of optimization problems into the meta-learner (Ravi & Larochelle, 2017; Finn et al., 2017a; Husken & Goerick, 2000). In particular, model-agnostic meta-learning (MAML) optimizes only for the initial parameters of the learner model, using standard gradient descent as the learner’s update rule (Finn et al., 2017a). Then, at meta-test time, the learner is trained via gradient descent. By incorporating prior knowledge about gradient-based learning, MAML improves on the statistical efficiency of black-box meta-learners and has successfully been applied to a range of meta-learning problems (Finn et al., 2017a; b; Li et al., 2017). But, does it do so at a cost? A natural question that arises with purely gradient-based meta-learners such as MAML is whether it is indeed sufficient to only learn an initialization, or whether representational power is in fact lost from not learning the update rule. Intuitively, we might surmise that learning an update rule is more expressive than simply learning an initialization for gradient descent. In this paper, we seek to answer the following question: does simply learning the initial parameters of a deep neural network have the same representational power as arbitrarily expressive meta-learners that directly ingest the training data at meta-test time? Or, more concisely, does representation combined with standard gradient descent have sufficient capacity to constitute any learning algorithm?

We analyze this question from the standpoint of the universal function approximation theorem. We compare the theoretical representational capacity of the two meta-learning approaches: a deep network updated with one gradient step, and a meta-learner that directly ingests a training set and test input and outputs predictions for that test input (e.g. using a recurrent neural network). In studying the universality of MAML, we find that, for a sufficiently deep learner model, MAML has the same theoretical representational power as recurrent meta-learners. We therefore conclude that, when using deep, expressive function approximators, there is no theoretical disadvantage in terms of representational power to using MAML over a black-box meta-learner represented, for example, by a recurrent network.

Since MAML has the same representational power as any other universal meta-learner, the next question we might ask is: what is the benefit of using MAML over any other approach? We study this question by analyzing the effect of continuing optimization on MAML performance. Although MAML optimizes a network’s parameters for maximal performance after a fixed small number of gradient steps, we analyze the effect of taking substantially more gradient steps at meta-test time. We find that initializations learned by MAML are extremely resilient to overfitting to tiny datasets, in stark contrast to more conventional network initialization, even when taking many more gradient steps than were used during meta-training. We also find that the MAML initialization is substantially better suited for extrapolation beyond the distribution of tasks seen at meta-training time, when compared to meta-learning methods based on networks that ingest the entire training set. We analyze this setting empirically and provide some intuition to explain this effect.

Preliminaries

In this section, we review the universal function approximation theorem and its extensions that we will use when considering the universal approximation of learning algorithms. We also overview the model-agnostic meta-learning algorithm and an architectural extension that we will use in Section 4.

2 Model-Agnostic Meta-Learning with a Bias Transformation

Model-Agnostic Meta-Learning (MAML) is a method that proposes to learn an initial set of parameters θ\theta such that one or a few gradient steps on θ\theta computed using a small amount of data for one task leads to effective generalization on that task (Finn et al., 2017a). Tasks typically correspond to supervised classification or regression problems, but can also correspond to reinforcement learning problems. The MAML objective is computed over many tasks {Tj}\{\mathcal{T}_{j}\} as follows:

where DTj\mathcal{D}_{\mathcal{T}_{j}} corresponds to a training set for task Tj\mathcal{T}_{j} and the outer loss evaluates generalization on test data in DTj′\mathcal{D}^{\prime}_{\mathcal{T}_{j}}. The inner optimization to compute θTj′\theta_{\mathcal{T}_{j}}^{\prime} can use multiple gradient steps; though, in this paper, we will focus on the single gradient step setting. After meta-training on a wide range of tasks, the model can quickly and efficiently learn new, held-out test tasks by running gradient descent starting from the meta-learned representation θ\theta.

While MAML is compatible with any neural network architecture and any differentiable loss function, recent work has observed that some architectural choices can improve its performance. A particularly effective modification, introduced by Finn et al. (2017b), is to concatenate a vector of parameters, θb\theta_{b}, to the input. As with all other model parameters, θb\theta_{b} is updated in the inner loop via gradient descent, and the initial value of θb\theta_{b} is meta-learned. This modification, referred to as a bias transformation, increases the expressive power of the error gradient without changing the expressivity of the model itself. While Finn et al. (2017b) report empirical benefit from this modification, we will use this architectural design as a symmetry-breaking mechanism in our universality proof.

Meta-Learning and Universality

We can broadly classify RNN-based meta-learning methods into two categories. In the first approach (Santoro et al., 2016; Duan et al., 2016; Wang et al., 2016; Mishra et al., 2018), there is a meta-learner model gg with parameters ϕ\phi which takes as input the dataset DT\mathcal{D}_{\mathcal{T}} for a particular task T\mathcal{T} and a new test input x⋆\mathbf{x}^{\star}, and outputs the estimated output y^⋆\hat{\mathbf{y}}^{\star} for that input:

The meta-learner gg is typically a recurrent model that iterates over the dataset D\mathcal{D} and the new input x⋆\mathbf{x}^{\star}. For a recurrent neural network model that satisfies the UFA theorem, this approach is maximally expressive, as it can represent any function on the dataset DT\mathcal{D}_{\mathcal{T}} and test input x⋆\mathbf{x}^{\star}.

In the second approach (Hochreiter et al., 2001; Bengio et al., 1992; Li & Malik, 2017b; Andrychowicz et al., 2016; Ravi & Larochelle, 2017; Ha et al., 2017), there is a meta-learner gg that takes as input the dataset for a particular task DT\mathcal{D}_{\mathcal{T}} and the current weights θ\theta of a learner model ff, and outputs new parameters θT′\theta_{\mathcal{T}}^{\prime} for the learner model. Then, the test input x⋆\mathbf{x}^{\star} is fed into the learner model to produce the predicted output y^⋆\hat{\mathbf{y}}^{\star}. The process can be written as follows:

Note that, in the form written above, this approach can be as expressive as the previous approach, since the meta-learner could simply copy the dataset into some of the predicted weights, reducing to a model that takes as input the dataset and the test example.For this to be possible, the model ff must be a neural network with at least two hidden layers, since the dataset can be copied into the first layer of weights and the predicted output must be a universal function approximator of both the dataset and the test input. Several versions of this approach, i.e. Ravi & Larochelle (2017); Li & Malik (2017b), have the recurrent meta-learner operate on order-invariant features such as the gradient and objective value averaged over the datapoints in the dataset, rather than operating on the individual datapoints themselves. This induces a potentially helpful inductive bias that disallows coupling between datapoints, ignoring the ordering within the dataset. As a result, the meta-learning process can only produce permutation-invariant functions of the dataset.

In model-agnostic meta-learning (MAML), instead of using an RNN to update the weights of the learner ff, standard gradient descent is used. Specifically, the prediction y^⋆\hat{\mathbf{y}}^{\star} for a test input x⋆\mathbf{x}^{\star} is:

The first goal of this paper is to show that fMAML(DT,x⋆;θ)f_{\text{MAML}}(\mathcal{D}_{\mathcal{T}},\mathbf{x}^{\star};\theta) is a universal function approximator of (DT,x⋆)(\mathcal{D}_{\mathcal{T}},\mathbf{x}^{\star}) in the one-shot setting, where the dataset DT\mathcal{D}_{\mathcal{T}} consists of a single datapoint (x,y)(\mathbf{x},\mathbf{y}). Then, we will consider the case of KK-shot learning, showing that fMAML(DT,x⋆;θ)f_{\text{MAML}}(\mathcal{D}_{\mathcal{T}},\mathbf{x}^{\star};\theta) is universal in the set of functions that are invariant to the permutation of datapoints. In both cases, we will discuss meta supervised learning problems with both discrete and continuous labels and the loss functions under which universality does or does not hold.

Universality of the One-Shot Gradient-Based Learner

We first introduce a proof of the universality of gradient-based meta-learning for the special case with only one training point, corresponding to one-shot learning. We denote the training datapoint as (x,y)(\mathbf{x},\mathbf{y}), and the test input as x⋆\mathbf{x}^{\star}. A universal learning algorithm approximator corresponds to the ability of a meta-learner to represent any function ftarget(x,y,x⋆)f_{\text{target}}(\mathbf{x},\mathbf{y},\mathbf{x}^{\star}) up to arbitrary precision.

where ϕ(⋅;θft,θb)\phi(\cdot;\theta_{\text{ft}},\theta_{b}) represents an input feature extractor with parameters θft\theta_{\text{ft}} and a scalar bias transformation variable θb\theta_{b}, ∏i=1NWi\prod_{i=1}^{N}W_{i} is a product of square linear weight matrices, fout(⋅,θout)f_{\text{out}}(\cdot,\theta_{\text{out}}) is a function at the output, and the learned parameters are θ:={θft,θb,{Wi},θout}\theta:=\{\theta_{\text{ft}},\theta_{b},\{W_{i}\},\theta_{\text{out}}\}. The input feature extractor and output function can be represented with fully connected neural networks with one or more hidden layers, which we know are universal function approximators, while ∏i=1NWi\prod_{i=1}^{N}W_{i} corresponds to a set of linear layers with non-negative input and activations.

where we will disregard the last term, assuming that α\alpha is comparatively small such that α2\alpha^{2} and all higher order terms vanish. In general, these terms do not necessarily need to vanish, and likely would further improve the expressiveness of the gradient update, but we disregard them here for the sake of the simplicity of the derivation. Ignoring these terms, we now note that the post-update value of z⋆\mathbf{z}^{\star} when x⋆\mathbf{x}^{\star} is provided as input into f^(⋅;θ′)\hat{f}(\cdot;\theta^{\prime}) is given by

and f^(x⋆;θ′)=fout(z⋆;θout′)\hat{f}(\mathbf{x}^{\star};\theta^{\prime})=f_{\text{out}}(\mathbf{z}^{\star};\theta_{\text{out}}^{\prime}).

Our goal is to show that that there exists a setting of WiW_{i}, foutf_{\text{out}}, and ϕ\phi for which the above function, f^(x⋆,θ′)\hat{f}(\mathbf{x}^{\star},\theta^{\prime}), can approximate any function of (x,y,x⋆)(\mathbf{x},\mathbf{y},\mathbf{x}^{\star}). To show universality, we will aim to independently control information flow from x\mathbf{x}, from y\mathbf{y}, and from x⋆\mathbf{x}^{\star} by multiplexing forward information from x\mathbf{x} and backward information from y\mathbf{y}. We will achieve this by decomposing WiW_{i}, ϕ\phi, and the error gradient into three parts, as follows:

where A1=IA_{1}=I, BN=IB_{N}=I, AiA_{i} can be chosen to be any symmetric positive-definite matrix, and BiB_{i} can be chosen to be any positive definite matrix. In Appendix D, we further show that these definitions of the weight matrices satisfy the condition that the activations are non-negative, meaning that the model f^\hat{f} can be represented by a generic deep network with ReLU nonlinearities.

Finally, we need to define the function foutf_{\text{out}} at the output. When the training input x\mathbf{x} is passed in, we need foutf_{\text{out}} to propagate information about the label y\mathbf{y} as defined in Equation 2. And, when the test input x⋆\mathbf{x}^{\star} is passed in, we need a different function defined only on z‾⋆\overline{\mathbf{z}}^{\star}. Thus, we will define foutf_{\text{out}} as a neural network that approximates the following multiplexer function and its derivatives (as shown possible by Hornik et al. (1990)):

Now, combining Equations 3 and 5, we can see that the post-update value is the following:

Let us assume that e‾(y)\overline{e}(\mathbf{y}) can be chosen to be any linear (but not affine) function of y\mathbf{y}. Then, we can choose θft\theta_{\text{ft}}, θh\theta_{h}, {Ai;i>1}\{A_{i};i>1\}, {Bi;i<N}\{B_{i};i<N\} such that the function

Intuitively, Equation 7 can be viewed as a sum of basis vectors Aie‾(y)A_{i}\overline{e}(\mathbf{y}) weighted by ki(x,x⋆)k_{i}(\mathbf{x},\mathbf{x}^{\star}), which is passed into hposth_{\text{post}} to produce the output. There are likely a number of ways to prove Lemma 4.1. In Appendix A.1, we provide a simple though inefficient proof, which we will briefly summarize here. We can define kik_{i} to be a indicator function, indicating when (x,x⋆)(\mathbf{x},\mathbf{x}^{\star}) takes on a particular value indexed by ii. Then, we can define Aie‾(y)A_{i}\overline{e}(\mathbf{y}) to be a vector containing the information of y\mathbf{y} and ii. Then, the result of the summation will be a vector containing information about the label y\mathbf{y} and the value of (x,x⋆)(\mathbf{x},\mathbf{x}^{\star}) which is indexed by ii. Finally, hposth_{\text{post}} defines the output for each value of (x,y,x⋆)(\mathbf{x},\mathbf{y},\mathbf{x}^{\star}). The bias transformation variable θb\theta_{b} plays a vital role in our construction, as it breaks the symmetry within ki(x,x⋆)k_{i}(\mathbf{x},\mathbf{x}^{\star}). Without such asymmetry, it would not be possible for our constructed function to represent any function of x\mathbf{x} and x⋆\mathbf{x}^{\star} after one gradient step.

In conclusion, we have shown that there exists a neural network structure for which f^(x⋆;θ′)\hat{f}(\mathbf{x}^{\star};\theta^{\prime}) is a universal approximator of ftarget(x,y,x⋆)f_{\text{target}}(\mathbf{x},\mathbf{y},\mathbf{x}^{\star}). We chose a particular form of f^(⋅;θ)\hat{f}(\cdot;\theta) that decouples forward and backward information flow. With this choice, it is possible to impose any desired post-update function, even in the face of adversarial training datasets and loss functions, e.g. when the gradient points in the wrong direction. If we make the assumption that the inner loss function and training dataset are not chosen adversarially and the error gradient points in the direction of improvement, it is likely that a much simpler architecture will suffice that does not require multiplexing of forward and backward information in separate channels. Informative loss functions and training data allowing for simpler functions is indicative of the inductive bias built into gradient-based meta-learners, which is not present in recurrent meta-learners.

Our result in this section implies that a sufficiently deep representation combined with just a single gradient step can approximate any one-shot learning algorithm. In the next section, we will show the universality of MAML for KK-shot learning algorithms.

General Universality of the Gradient-Based Learner

Now, we consider the more general KK-shot setting, aiming to show that MAML can approximate any permutation invariant function of a dataset and test datapoint ({(x,y)i;i∈1...K},x⋆)(\{(\mathbf{x},\mathbf{y})_{i};i\in 1...K\},\mathbf{x}^{\star}) for K>1K>1. Note that KK does not need to be small. To reduce redundancy, we will only overview the differences from the 11-shot setting in this section. We include a full proof in Appendix B.

In the KK-shot setting, the parameters of f^(⋅,θ)\hat{f}(\cdot,\theta) are updated according to the following rule:

Defining the form of f^\hat{f} to be the same as in Section 4, the post-update function is the following:

Loss Functions

In the previous sections, we showed that a deep representation combined with gradient descent can approximate any learning algorithm. In this section, we will discuss the requirements that the loss function must satisfy in order for the results in Sections 4 and 5 to hold. As one might expect, the main requirement will be for the label to be recoverable from the gradient of the loss.

As seen in the definition of foutf_{\text{out}} in Equation 4, the pre-update function f^(x,θ)\hat{f}(\mathbf{x},\theta) is given by gpre(z;θg)g_{\text{pre}}(\mathbf{z};\theta_{g}), where gpreg_{\text{pre}} is used for back-propagating information about the label(s) to the learner. As stated in Equation 2, we require that the error gradient with respect to z\mathbf{z} to be:

and where e‾(y)\overline{e}(\mathbf{y}) and eˇ(y)\check{e}(\mathbf{y}) must be able to represent [at least] any linear function of the label y\mathbf{y}.

The gradient of the standard mean-squared error objective evaluated at y^=0\hat{\mathbf{y}}=\mathbf{0} is a linear, invertible function of y\mathbf{y}.

The gradient of the softmax cross entropy loss with respect to the pre-softmax logits is a linear, invertible function of y\mathbf{y}, when evaluated at 0\mathbf{0}.

Experiments

Now that we have shown that meta-learners that use standard gradient descent with a sufficiently deep representation can approximate any learning procedure, and are equally expressive as recurrent learners, a natural next question is – is there empirical benefit to using one meta-learning approach versus another, and in which cases? To answer this question, we next aim to empirically study the inductive bias of gradient-based and recurrent meta-learners. Then, in Section 7.2, we will investigate the role of model depth in gradient-based meta-learning, as the theory suggests that deeper networks lead to increased expressive power for representing different learning procedures.

First, we aim to empirically explore the differences between gradient-based and recurrent meta-learners. In particular, we aim to answer the following questions: (1) can a learner trained with MAML further improve from additional gradient steps when learning new tasks at test time, or does it start to overfit? and (2) does the inductive bias of gradient descent enable better few-shot learning performance on tasks outside of the training distribution, compared to learning algorithms represented as recurrent networks?

To study both questions, we will consider two simple few-shot learning domains. The first is 5-shot regression on a family of sine curves with varying amplitude and phase. We trained all models on a uniform distribution of tasks with amplitudes A∈[0.1,5.0]A\in[0.1,5.0], and phases γ∈[0,π]\gamma\in[0,\pi]. The second domain is 1-shot character classification using the Omniglot dataset (Lake et al., 2011), following the training protocol introduced by Santoro et al. (2016). In our comparisons to recurrent meta-learners, we will use two state-of-the-art meta-learning models: SNAIL (Mishra et al., 2018) and meta-networks (Munkhdalai & Yu, 2017). In some experiments, we will also compare to a task-conditioned model, which is trained to map from both the input and the task description to the label. Like MAML, the task-conditioned model can be fine-tuned on new data using gradient descent, but is not trained for few-shot adaptation. We include more experimental details in Appendix G.

To answer the first question, we fine-tuned a model trained using MAML with many more gradient steps than used during meta-training. The results on the sinusoid domain, shown in Figure 2, show that a MAML-learned initialization trained for fast adaption in 5 steps can further improve beyond 55 gradient steps, especially on out-of-distribution tasks. In contrast, a task-conditioned model trained without MAML can easily overfit to out-of-distribution tasks. With the Omniglot dataset, as seen in Figure 4, a MAML model that was trained with 55 inner gradient steps can be fine-tuned for 100100 gradient steps without leading to any drop in test accuracy. As expected, a model initialized randomly and trained from scratch quickly reaches perfect training accuracy, but overfits massively to the 2020 examples.

Next, we investigate the second question, aiming to compare MAML with state-of-the-art recurrent meta-learners on tasks that are related to, but outside of the distribution of the training tasks. All three methods achieved similar performance within the distribution of training tasks for 5-way 1-shot Omniglot classification and 5-shot sinusoid regression. In the Omniglot setting, we compare each method’s ability to distinguish digits that have been sheared or scaled by varying amounts. In the sinusoid regression setting, we compare on sinusoids with extrapolated amplitudes within [5.0,10.0][5.0,10.0] and phases within [π,2π][\pi,2\pi]. The results in Figure 3 and Appendix G show a clear trend that MAML recovers more generalizable learning strategies. Combined with the theoretical universality results, these experiments indicate that deep gradient-based meta-learners are not only equivalent in representational power to recurrent meta-learners, but should also be a considered as a strong contender in settings that contain domain shift between meta-training and meta-testing tasks, where their strong inductive bias for reasonable learning strategies provides substantially improved performance.

2 Effect of Depth

The proofs in Sections 4 and 5 suggest that gradient descent with deeper representations results in more expressive learning procedures. In contrast, the universal function approximation theorem only requires a single hidden layer to approximate any function. Now, we seek to empirically explore this theoretical finding, aiming to answer the question: is there a scenario for which model-agnostic meta-learning requires a deeper representation to achieve good performance, compared to the depth of the representation needed to solve the underlying tasks being learned?

To answer this question, we will study a simple regression problem, where the meta-learning goal is to infer a polynomial function from 40 input/output datapoints. We use polynomials of degree 33 where the coefficients and bias are sampled uniformly at random within andtheinputvaluesrangewithinand the input values range within. Similar to the conditions in the proof, we meta-train and meta-test with one gradient step, use a mean-squared error objective, use ReLU nonlinearities, and use a bias transformation variable of dimension 10. To compare the relationship between depth and expressive power, we will compare models with a fixed number of parameters, approximately 40,00040,000, and vary the network depth from 1 to 5 hidden layers. As a point of comparison to the models trained for meta-learning using MAML, we trained standard feedforward models to regress from the input and the 4-dimensional task description (the 3 coefficients of the polynomial and the scalar bias) to the output. These task-conditioned models act as an oracle and are meant to empirically determine the depth needed to represent these polynomials, independent of the meta-learning process. Theoretically, we would expect the task-conditioned models to require only one hidden layer, as per the universal function approximation theorem. In contrast, we would expect the MAML model to require more depth. The results, shown in Figure 5, demonstrate that the task-conditioned model does indeed not benefit from having more than one hidden layer, whereas the MAML clearly achieves better performance with more depth even though the model capacity, in terms of the number of parameters, is fixed. This empirical effect supports the theoretical finding that depth is important for effective meta-learning using MAML.

Conclusion

In this paper, we show that there exists a form of deep neural network such that the initial weights combined with gradient descent can approximate any learning algorithm. Our findings suggest that, from the standpoint of expressivity, there is no theoretical disadvantage to embedding gradient descent into the meta-learning process. In fact, in all of our experiments, we found that the learning strategies acquired with MAML are more successful when faced with out-of-domain tasks compared to recurrent learners. Furthermore, we show that the representations acquired with MAML are highly resilient to overfitting. These results suggest that gradient-based meta-learning has a number of practical benefits, and no theoretical downsides in terms of expressivity when compared to alternative meta-learning models. Independent of the type of meta-learning algorithm, we formalize what it means for a meta-learner to be able to approximate any learning algorithm in terms of its ability to represent functions of the dataset and test inputs. This formalism provides a new perspective on the learning-to-learn problem, which we hope will lead to further discussion and research on the goals and methodology surrounding meta-learning.

We thank Sharad Vikram for detailed feedback on the proof, as well as Justin Fu, Ashvin Nair, and Kelvin Xu for feedback on an early draft of this paper. We also thank Erin Grant for helpful conversations and Nikhil Mishra for providing code for SNAIL. This research was supported by the National Science Foundation through IIS-1651843 and a Graduate Research Fellowship, as well as NVIDIA.

References

Appendix A Supplementary Proofs for 1-Shot Setting

While there are likely a number of ways to prove Lemma 4.1 (copied below for convenience), here we provide a simple, though inefficient, proof of Lemma 4.1. See 4.1

To prove this lemma, we will proceed by showing that we can choose e‾\overline{e}, θft\theta_{\text{ft}}, and each AiA_{i} and BiB_{i} such that the summation contains a complete description of the values of x\mathbf{x}, x⋆\mathbf{x}^{\star}, and y\mathbf{y}. Then, because hposth_{\text{post}} is a universal function approximator, f^(x⋆,θ′)\hat{f}(\mathbf{x}^{\star},\theta^{\prime}) will be able to approximate any function of x\mathbf{x}, x⋆\mathbf{x}^{\star}, and y\mathbf{y}.

Since A1=IA_{1}=I and BN=IB_{N}=I, we will essentially ignore the first and last elements of the sum by defining B1:=ϵIB_{1}:=\epsilon I and AN:=ϵIA_{N}:=\epsilon I, where ϵ\epsilon is a small positive constant to ensure positive definiteness. Then, we can rewrite the summation, omitting the first and last terms:

Next, we will re-index using two indexing variables, jj and ll, where jj will index over the discretization of x\mathbf{x} and ll over the discretization of x⋆\mathbf{x}^{\star}.

Next, we will define our chosen form of kjlk_{jl} in Equation 8. We show how to acquire this form in the next section.

We can choose θft\theta_{\text{ft}} and each BjlB_{jl} such that

where discr(⋅)\textnormal{discr}(\cdot) denotes a function that produces a one-hot discretization of its input and e\mathbf{e} denotes the 0-indexed standard basis vector.

Now that we have defined the function kjlk_{jl}, we will next define the other terms in the sum. Our goal is for the summation to contain complete information about (x,x⋆,y)(\mathbf{x},\mathbf{x}^{\star},\mathbf{y}). To do so, we will chose e‾(y)\overline{e}(\mathbf{y}) to be the linear function that outputs J∗LJ*L stacked copies of y\mathbf{y}. Then, we will define AjlA_{jl} to be a matrix that selects the copy of y\mathbf{y} in the position corresponding to (j,l)(j,l), i.e. in the position j+J∗lj+J*l. This can be achieved using a diagonal AjlA_{jl} matrix with diagonal values of 1+ϵ1+\epsilon at the positions corresponding to the kkth vector, and ϵ\epsilon elsewhere, where k=(j+J∗l)k=(j+J*l) and ϵ\epsilon is used to ensure that AjlA_{jl} is positive definite.

As a result, the post-update function is as follows:

where y\mathbf{y} is at the position j+J∗lj+J*l within the vector v(x,x⋆,y)v(\mathbf{x},\mathbf{x}^{\star},\mathbf{y}), where jj satisfies discr(x)=ej\text{discr}(\mathbf{x})=\mathbf{e}_{j} and where ll satisfies discr(x⋆)=el\text{discr}(\mathbf{x}^{\star})=\mathbf{e}_{l}. Note that the vector −αv(x,x⋆,y)-\alpha v(\mathbf{x},\mathbf{x}^{\star},\mathbf{y}) is a complete description of (x,x⋆,y)(\mathbf{x},\mathbf{x}^{\star},\mathbf{y}) in that x\mathbf{x}, x⋆\mathbf{x}^{\star}, and y\mathbf{y} can be decoded from it. Therefore, since hposth_{\text{post}} is a universal function approximator and because its input contains all of the information of (x,x⋆,y)(\mathbf{x},\mathbf{x}^{\star},\mathbf{y}), the function f^(x⋆;θ′)≈hpost(−αv(x,x⋆,y);θh)\hat{f}(\mathbf{x}^{\star};\theta^{\prime})\approx h_{\text{post}}\left(-\alpha v(\mathbf{x},\mathbf{x}^{\star},\mathbf{y});\theta_{h}\right) is a universal function approximator with respect to its inputs (x,x⋆,y)(\mathbf{x},\mathbf{x}^{\star},\mathbf{y}).

A.2 Proof of Lemma A.1

where we use EikE_{ik} to denote the matrix with a 1 at (i,k)(i,k) and 0 elsewhere, and ϵI\epsilon I is added to ensure the positive definiteness of BjlB_{jl} as required in the construction.

Using the above definitions, we can see that:

A.3 Form of linear weight matrices

Recall that we decomposed WiW_{i}, ϕ\phi, and the error gradient into three parts, as follows:

A.4 Output function

In this section, we will derive the post-update version of the output function fout(⋅;θout)f_{\text{out}}(\cdot;\theta_{\text{out}}). Recall that foutf_{\text{out}} is defined as a neural network that approximates the following multiplexer function and its derivatives (as shown possible by Hornik et al. (1990)):

The parameters {θg,θh}\{\theta_{g},\theta_{h}\} are a part of θout\theta_{\text{out}}, in addition to the parameters required to estimate the indicator functions and their corresponding products. Since z‾=0\overline{\mathbf{z}}=\mathbf{0} and hpost(z‾)=0h_{\text{post}}(\overline{\mathbf{z}})=\mathbf{0} when the gradient step is taken, we can see that the error gradients with respect to the parameters in the last term in Equation 12 will be approximately zero. Furthermore, as seen in the definition of gpreg_{\text{pre}} in Section 6, the value of gpre(z,θg)g_{\text{pre}}(\mathbf{z},\theta_{g}) is also zero, resulting in a gradient of approximately zero for the first indicator function.To guarantee that g and h are zero when evaluated at x\mathbf{x}, we make the assumption that gpreg_{\text{pre}} and hposth_{\text{post}} are neural networks with no biases and nonlinearity functions that output zero when evaluated at zero.

The post-update value of foutf_{\text{out}} is therefore:

as long as z‾⋆≠0\overline{\mathbf{z}}^{\star}\neq\mathbf{0}. In Appendix A.1, we can see that z⋆\mathbf{z}^{\star} is indeed not equal to zero.

Appendix B Full K-Shot Proof of Universality

In this appendix, we provide a full proof of the universality of gradient-based meta-learning in the general case with K>1K>1 datapoints. This proof will share a lot of content from the proof in the 11-shot setting, but we include it for completeness.

We aim to show that a deep representation combined with one step of gradient descent can approximate any permutation invariant function of a dataset and test datapoint ({(x,y)i;i∈1...K},x⋆)(\{(\mathbf{x},\mathbf{y})_{i};i\in 1...K\},\mathbf{x}^{\star}) for K>1K>1. Note that KK does not need to be small.

We will start by constructing f^\hat{f}. With the same motivation as in Section 4, we will construct f^(⋅;θ)\hat{f}(\cdot;\theta) as the following:

ϕ(⋅;θft,θb)\phi(\cdot;\theta_{\text{ft}},\theta_{b}) represents an input feature extractor with parameters θft\theta_{\text{ft}} and a scalar bias transformation variable θb\theta_{b}, ∏i=1NWi\prod_{i=1}^{N}W_{i} is a product of square linear weight matrices, fout(⋅,θout)f_{\text{out}}(\cdot,\theta_{\text{out}}) is a readout function at the output, and the learned parameters are θ:={θft,θb,{Wi},θout}\theta:=\{\theta_{\text{ft}},\theta_{b},\{W_{i}\},\theta_{\text{out}}\}. The input feature extractor and readout function can be represented with fully connected neural networks with one or more hidden layers, which we know are universal function approximators, while ∏i=1NWi\prod_{i=1}^{N}W_{i} corresponds to a set of linear layers. Note that deep ReLU networks act like deep linear networks when the input and pre-synaptic activations are non-negative. We will later show that this is indeed the case within these linear layers, meaning that the neural network function f^\hat{f} is fully generic and can be represented by deep ReLU networks, as visualized in Figure 1.

Therefore, the post-update value of ∏i=1NWi′=∏i=1N(Wi−α1K∑k∇Wi)\prod_{i=1}^{N}W_{i}^{\prime}=\prod_{i=1}^{N}(W_{i}-\alpha\frac{1}{K}\sum_{k}\nabla_{W_{i}}) is given by

where we move the summation over kk to the left and where we will disregard the last term, assuming that α\alpha is comparatively small such that α2\alpha^{2} and all higher order terms vanish. In general, these terms do not necessarily need to vanish, and likely would further improve the expressiveness of the gradient update, but we disregard them here for the sake of the simplicity of the derivation. Ignoring these terms, we now note that the post-update value of z⋆\mathbf{z}^{\star} when x⋆\mathbf{x}^{\star} is provided as input into f^(⋅;θ′)\hat{f}(\cdot;\theta^{\prime}) is given by

and f^(x⋆;θ′)=fout(z⋆;θout′)\hat{f}(\mathbf{x}^{\star};\theta^{\prime})=f_{\text{out}}(\mathbf{z}^{\star};\theta_{\text{out}}^{\prime}).

Our goal is to show that that there exists a setting of WiW_{i}, foutf_{\text{out}}, and ϕ\phi for which the above function, f^(x⋆,θ′)\hat{f}(\mathbf{x}^{\star},\theta^{\prime}), can approximate any function of ({(x,y)k},x⋆)(\{(\mathbf{x},\mathbf{y})_{k}\},\mathbf{x}^{\star}). To show universality, we will aim independently control information flow from {xk}\{\mathbf{x}_{k}\}, from {yk}\{\mathbf{y}_{k}\}, and from x⋆\mathbf{x}^{\star} by multiplexing forward information from {xk}\{\mathbf{x}_{k}\} and x⋆\mathbf{x}^{\star} and backward information from {yk}\{\mathbf{y}_{k}\}. We will achieve this by decomposing WiW_{i}, ϕ\phi, and the error gradient into three parts, as follows:

where A1=IA_{1}=I, BN=IB_{N}=I, AiA_{i} can be chosen to be any symmetric positive-definite matrix, and BiB_{i} can be chosen to be any positive definite matrix. In Appendix D, we will further show that these definitions of the weight matrices satisfy the condition that their activations are non-negative, meaning that the model f^\hat{f} can be represented by a generic deep network with ReLU nonlinearities.

Finally, we need to define the function foutf_{\text{out}} at the output. When a training input xk\mathbf{x}_{k} is passed in, we need foutf_{\text{out}} to propagate information about its corresponding label yk\mathbf{y}_{k} as defined in Equation 15. And, when the test input x⋆\mathbf{x}^{\star} is passed in, we need a function defined on z‾⋆\overline{\mathbf{z}}^{\star}. Thus, we will define foutf_{\text{out}} as a neural network that approximates the following multiplexer function and its derivatives (as shown possible by Hornik et al. (1990)):

Now, combining Equations 16 and 18, we can see that the post-update value is the following:

In conclusion, we have shown that there exists a neural network structure for which f^(x⋆;θ′)\hat{f}(\mathbf{x}^{\star};\theta^{\prime}) is a universal approximator of ftarget({(x,y)k},x⋆)f_{\text{target}}(\{(\mathbf{x},\mathbf{y})_{k}\},\mathbf{x}^{\star}).

Appendix C Supplementary Proof for K-Shot Setting

In Section 5 and Appendix B, we showed that the post-update function f^(x⋆;θ′)\hat{f}(\mathbf{x}^{\star};\theta^{\prime}) takes the following form:

In this section, we aim to show that the above form of f^(x⋆;θ′)\hat{f}(\mathbf{x}^{\star};\theta^{\prime}) can approximate any function of {(x,y)k;k∈1...K}\{(\mathbf{x},\mathbf{y})_{k};k\in 1...K\} and x⋆\mathbf{x}^{\star} that is invariant to the ordering of the training datapoints {(x,y)k;k∈1...K}\{(\mathbf{x},\mathbf{y})_{k};k\in 1...K\}. The proof will be very similar to the one-shot setting proof in Appendix A.1

Similar to Appendix A.1, we will ignore the first and last elements of the sum by defining B1B_{1} to be ϵI\epsilon I and ANA_{N} to be ϵI\epsilon I, where ϵ\epsilon is a small positive constant to ensure positive definiteness. We will then re-index the first summation over i=2...N−1i=2...N-1 to instead use two indexing variables jj and ll as follows:

As in Appendix A.1, we will define the function kjlk_{jl} to be an indicator function over the values of xk\mathbf{x}_{k} and x⋆\mathbf{x}^{\star}. In particular, we will reuse Lemma A.1, which was proved in Appendix A.2 and is copied below: See A.1

Likewise, we will chose e‾(yk)\overline{e}(\mathbf{y}_{k}) to be the linear function that outputs J∗LJ*L stacked copies of yk\mathbf{y}_{k}. Then, we will define AjlA_{jl} to be a matrix that selects the copy of yk\mathbf{y}_{k} in the position corresponding to (j,l)(j,l), i.e. in the position j+J∗lj+J*l. This can be achieved using a diagonal AjlA_{jl} matrix with diagonal values of 1+ϵ1+\epsilon at the positions corresponding to the nnth vector, and ϵ\epsilon elsewhere, where n=(j+J∗l)n=(j+J*l) and ϵ\epsilon is used to ensure that AjlA_{jl} is positive definite.

As a result, the post-update function is as follows:

where yk\mathbf{y}_{k} is at the position j+J∗lj+J*l within the vector v(xk,x⋆,yk)v(\mathbf{x}_{k},\mathbf{x}^{\star},\mathbf{y}_{k}), where jj satisfies discr(xk)=ej\text{discr}(\mathbf{x}_{k})=\mathbf{e}_{j} and where ll satisfies discr(xk⋆)=el\text{discr}(\mathbf{x}^{\star}_{k})=\mathbf{e}_{l}.

For discrete, one-shot labels yk\mathbf{y}_{k}, the summation over vv amounts to frequency counts of the triplets (xk,x⋆,yk)(\mathbf{x}_{k},\mathbf{x}^{\star},\mathbf{y}_{k}). In the setting with continuous labels, we cannot attain frequency counts, as we do not have access to a discretized version of the label. Thus, we must make the assumption that no two datapoints share the same input value xk\mathbf{x}_{k}. With this assumption, the summation over vv will contain the output values yk′\mathbf{y}_{k^{\prime}} at the index corresponding to the value of (xk′,x⋆)(\mathbf{x}_{k^{\prime}},\mathbf{x}^{\star}). For both discrete and continuous labels, this representation is redundant in x⋆\mathbf{x}^{\star}, but nonetheless contains sufficient information to decode the test input x⋆\mathbf{x}^{\star} and set of datapoints {(x,y)k}\{(\mathbf{x},\mathbf{y})_{k}\} (but not the order of datapoints).

Since hposth_{\text{post}} is a universal function approximator and because its input contains all of the information of ({(x,y)k},x⋆)(\{(\mathbf{x},\mathbf{y})_{k}\},\mathbf{x}^{\star}), the function f^(x⋆;θ′)≈hpost(−α1K∑k=1Kv(xk,x⋆,yk);θh)\hat{f}(\mathbf{x}^{\star};\theta^{\prime})\approx h_{\text{post}}\left(-\alpha\frac{1}{K}\sum_{k=1}^{K}v(\mathbf{x}_{k},\mathbf{x}^{\star},\mathbf{y}_{k});\theta_{h}\right) is a universal function approximator with respect to {(x,y)k}\{(\mathbf{x},\mathbf{y})_{k}\} and x⋆\mathbf{x}^{\star}.

Appendix D Deep ReLU Networks

In this appendix, we show that the network architecture with linear layers analyzed in the Sections 4 and 5 can be represented by a deep network with ReLU nonlinearities. We will do so by showing that the input and activations within the linear layers are all non-negative.

Appendix E Proof of Theorem 6.1

Here we provide a proof of Theorem 6.1: See 6.1

Appendix F Proof of Theorem 6.2

Here we provide a proof of Theorem 6.2: See 6.2

Appendix G Additional Experimental Details

In this section, we provide two additional comparisons on an out-of-distribution task and using additional gradient steps, shown in Figure 6. We also include additional experimental details.

For Omniglot, all meta-learning methods were trained using code provided by the authors of the respective papers, using the default model architectures and hyperparameters. The model embedding architecture was the same across all methods, using 4 convolutional layers with 3×33\times 3 kernels, 64 filters, stride 2, batch normalization, and ReLU nonlinearities. The convolutional layers were followed by a single linear layer. All methods used the Adam optimizer with default hyperparameters. Other hyperparameter choices were specific to the algorithm and can be found in the respective papers. For MAML in the sinusoid domain, we used a fully-connected network with two hidden layers of size 100, ReLU nonlinearities, and a bias transformation variable of size 10 concatenated to the input. This model was trained for 70,000 meta-iterations with 55 inner gradient steps of size α=0.001\alpha=0.001. For SNAIL in the sinusoid domain, the model consisted of 2 blocks of the following: 4 dilated convolutions with 2×12\times 1 kernels 16 channels, and dilation size of 1,2,4, and 8 respectively, then an attention block with key/value dimensionality of 8. The final layer is a 1×11\times 1 convolution to the output. Like MAML, this model was trained to convergence for 70,000 iterations using Adam with default hyperparameters. We evaluated the MAML and SNAIL models for 1200 trials, reporting the mean and 95%95\% confidence intervals. For computational reasons, we evaluated the MetaNet model using 600 trials, also reporting the mean and 95%95\% confidence intervals.

Following prior work (Santoro et al., 2016), we downsampled the Omniglot images to be 28×2828\times 28. When scaling or shearing the digits to produce out-of-domain data, we transformed the original 105×105105\times 105 Omniglot images, and then downsampled to 28×2828\times 28.

G.2 Depth Experiments

In the depth comparison, all models were trained to convergence using 70,000 iterations. Each model was defined to have a fixed number of hidden units based on the total number of parameters (fixed at around 40,000) and the number of hidden layers. Thus, the models with 22, 33, 44, and 55 hidden layers had 200200, 141141, 115115, and 100100 units per layer respectively. For the model with 11 hidden layer, we found that using more than 20,00020,000 hidden units, corresponding to 40,00040,000 parameters, resulted in poor performance. Thus, the results reported in the paper used a model with 11 hidden layer with 250250 units which performed much better. We trained each model three times and report the mean and standard deviation of the three runs. The performance of an individual run was computed using the average over