Meta-Learning with Implicit Gradients

Aravind Rajeswaran, Chelsea Finn, Sham Kakade, Sergey Levine

Introduction

A core aspect of intelligence is the ability to quickly learn new tasks by drawing upon prior experience from related tasks. Recent work has studied how meta-learning algorithms can acquire such a capability by learning to efficiently learn a range of tasks, thereby enabling learning of a new task with as little as a single example . Meta-learning algorithms can be framed in terms of recurrent or attention-based models that are trained via a meta-learning objective, to essentially encapsulate the learned learning procedure in the parameters of a neural network. An alternative formulation is to frame meta-learning as a bi-level optimization procedure , where the “inner” optimization represents adaptation to a given task, and the “outer” objective is the meta-training objective. Such a formulation can be used to learn the initial parameters of a model such that optimizing from this initialization leads to fast adaptation and generalization. In this work, we focus on this class of optimization-based methods, and in particular the model-agnostic meta-learning (MAML) formulation . MAML has been shown to be as expressive as black-box approaches , is applicable to a broad range of settings , and recovers a convergent and consistent optimization procedure .

Despite its appealing properties, meta-learning an initialization requires backpropagation through the inner optimization process. As a result, the meta-learning process requires higher-order derivatives, imposes a non-trivial computational and memory burden, and can suffer from vanishing gradients. These limitations make it harder to scale optimization-based meta learning methods to tasks involving medium or large datasets, or those that require many inner-loop optimization steps. Our goal is to develop an algorithm that addresses these limitations.

The main contribution of our work is the development of the implicit MAML (iMAML) algorithm, an approach for optimization-based meta-learning with deep neural networks that removes the need for differentiating through the optimization path. Our algorithm aims to learn a set of parameters such that an optimization algorithm that is initialized at and regularized to this parameter vector leads to good generalization for a variety of learning tasks. By leveraging the implicit differentiation approach, we derive an analytical expression for the meta (or outer level) gradient that depends only on the solution to the inner optimization and not the path taken by the inner optimization algorithm, as depicted in Figure 1. This decoupling of meta-gradient computation and choice of inner level optimizer has a number of appealing properties.

Problem Formulation and Notations

We first present the meta-learning problem in the context of few-shot supervised learning, and then generalize the notation to aid the rest of the exposition in the paper.

In the case of MAML , Alg(θ,D)\mathcal{A}lg({\bm{\theta}},\mathcal{D}) corresponds to one or multiple steps of gradient descent initialized at θ{\bm{\theta}}. For example, if one step of gradient descent is used, we have:

2 Proximal Regularization in the Inner Level

To have sufficient learning in the inner level while also avoiding over-fitting, Alg\mathcal{A}lg needs to incorporate some form of regularization. Since MAML uses a small number of gradient steps, this corresponds to early stopping and can be interpreted as a form of regularization and Bayesian prior . In cases like ill-conditioned optimization landscapes and medium-shot learning, we may want to take many gradient steps, which poses two challenges for MAML. First, we need to store and differentiate through the long optimization path of Alg\mathcal{A}lg, which imposes a considerable computation and memory burden. Second, the dependence of the model-parameters {ϕi}\{{\bm{\phi}}_{i}\} on the meta-parameters (θ)({\bm{\theta}}) shrinks and vanishes as the number of gradient steps in Alg\mathcal{A}lg grows, making meta-learning difficult. To overcome these limitations, we consider a more explicitly regularized algorithm:

3 The Bi-Level Optimization Problem

With this notation, the bi-level meta-learning problem can be written more generally as:

4 Total and Partial Derivatives

We use d\bm{d} to denote the total derivative and ∇\nabla to denote partial derivative. For nested function of the form Li(ϕi)\mathcal{L}_{i}({\bm{\phi}}_{i}) where ϕi=Algi(θ){\bm{\phi}}_{i}=\mathcal{A}lg_{i}({\bm{\theta}}), we have from chain rule

Note the important distinction between dθLi(Algi(θ))\bm{d}_{\bm{\theta}}\mathcal{L}_{i}(\mathcal{A}lg_{i}({\bm{\theta}})) and ∇ϕLi(Algi(θ))\nabla_{\bm{\phi}}\mathcal{L}_{i}(\mathcal{A}lg_{i}({\bm{\theta}})). The former passes derivatives through Algi(θ)\mathcal{A}lg_{i}({\bm{\theta}}) while the latter does not. ∇ϕLi(Algi(θ))\nabla_{\bm{\phi}}\mathcal{L}_{i}(\mathcal{A}lg_{i}({\bm{\theta}})) is simply the gradient function, i.e. ∇ϕLi(ϕ)\nabla_{\bm{\phi}}\mathcal{L}_{i}({\bm{\phi}}), evaluated at ϕ=Algi(θ){\bm{\phi}}=\mathcal{A}lg_{i}({\bm{\theta}}). Also note that dθLi(Algi(θ))\bm{d}_{\bm{\theta}}\mathcal{L}_{i}(\mathcal{A}lg_{i}({\bm{\theta}})) and ∇ϕLi(Algi(θ))\nabla_{\bm{\phi}}\mathcal{L}_{i}(\mathcal{A}lg_{i}({\bm{\theta}})) are dd–dimensional vectors, while dAlgi(θ)dθ\frac{d\mathcal{A}lg_{i}({\bm{\theta}})}{d{\bm{\theta}}} is a (d×d)(d\times d)–size Jacobian matrix. Throughout this text, we will also use dθ\bm{d}_{\bm{\theta}} and ddθ\frac{d}{d{\bm{\theta}}} interchangeably.

The Implicit MAML Algorithm

Our aim is to solve the bi-level meta-learning problem in Eq. LABEL:eq:update_rule using an iterative gradient based algorithm of the form θ←θ−η dθF(θ){\bm{\theta}}\leftarrow{\bm{\theta}}-\eta\ \bm{d}_{\bm{\theta}}F({\bm{\theta}}). Although we derive our method based on standard gradient descent for simplicity, any other optimization method, such as quasi-Newton or Newton methods, Adam , or gradient descent with momentum can also be used without modification. The gradient descent update be expanded using the chain rule as

Here, ∇ϕLi(Algi⋆(θ))\nabla_{\bm{\phi}}\mathcal{L}_{i}(\mathcal{A}lg^{\star}_{i}({\bm{\theta}})) is simply ∇ϕLi(ϕ)∣ϕ=Algi⋆(θ)\nabla_{\bm{\phi}}\mathcal{L}_{i}({\bm{\phi}})\mid_{{\bm{\phi}}=\mathcal{A}lg^{\star}_{i}({\bm{\theta}})} which can be easily obtained in practice via automatic differentiation. For this update rule, we must compute dAlgi⋆(θ)dθ\frac{d\mathcal{A}lg^{\star}_{i}({\bm{\theta}})}{d{\bm{\theta}}}, where Algi⋆\mathcal{A}lg^{\star}_{i} is implicitly defined as an optimization problem (Eq. LABEL:eq:update_rule), which presents the primary challenge. We now present an efficient algorithm (in compute and memory) to compute the meta-gradient..

If Algi⋆(θ)\mathcal{A}lg^{\star}_{i}({\bm{\theta}}) is implemented as an iterative algorithm, such as gradient descent, then one way to compute dAlgi⋆(θ)dθ\frac{d\mathcal{A}lg^{\star}_{i}({\bm{\theta}})}{d{\bm{\theta}}} is to propagate derivatives through the iterative process, either in forward mode or reverse mode. However, this has the drawback of depending explicitly on the path of the optimization, which has to be fully stored in memory, quickly becoming intractable when the number of gradient steps needed is large. Furthermore, for second order optimization methods, such as Newton’s method, third derivatives are needed which are difficult to obtain. Furthermore, this approach becomes impossible when non-differentiable operations, such as line-searches, are used. However, by recognizing that Algi⋆\mathcal{A}lg^{\star}_{i} is implicitly defined as the solution to an optimization problem, we may employ a different strategy that does not need to consider the path of the optimization but only the final result. This is derived in the following Lemma.

(Implicit Jacobian) Consider Algi⋆(θ)\mathcal{A}lg^{\star}_{i}({\bm{\theta}}) as defined in Eq. LABEL:eq:update_rule for task Ti\mathcal{T}_{i}. Let ϕi=Algi⋆(θ){\bm{\phi}}_{i}=\mathcal{A}lg^{\star}_{i}({\bm{\theta}}) be the result of Algi⋆(θ)\mathcal{A}lg^{\star}_{i}({\bm{\theta}}). If (I+1λ∇ϕ2L^i(ϕi))\left(\bm{I}+\frac{1}{\lambda}\nabla_{\bm{\phi}}^{2}\hat{\mathcal{L}}_{i}({\bm{\phi}}_{i})\right) is invertible, then the derivative Jacobian is

Note that the derivative (Jacobian) depends only on the final result of the algorithm, and not the path taken by the algorithm. Thus, in principle any approach of algorithm can be used to compute Algi⋆(θ)\mathcal{A}lg^{\star}_{i}({\bm{\theta}}), thereby decoupling meta-gradient computation from choice of inner level optimizer.

Practical Algorithm: While Lemma 1 provides an idealized way to compute the Algi⋆\mathcal{A}lg^{\star}_{i} Jacobians and thus by extension the meta-gradient, it may be difficult to directly use it in practice. Two issues are particularly relevant. First, the meta-gradients require computation of Algi⋆(θ)\mathcal{A}lg^{\star}_{i}({\bm{\theta}}), which is the exact solution to the inner optimization problem. In practice, we may be able to obtain only approximate solutions. Second, explicitly forming and inverting the matrix in Eq. 6 for computing the Jacobian may be intractable for large deep neural networks. To address these difficulties, we consider approximations to the idealized approach that enable a practical algorithm.

First, we consider an approximate solution to the inner optimization problem, that can be obtained with iterative optimization algorithms like gradient descent.

(δ\delta–approx. algorithm) Let Algi(θ)\mathcal{A}lg_{i}({\bm{\theta}}) be a δ\delta–accurate approximation of Algi⋆(θ)\mathcal{A}lg^{\star}_{i}({\bm{\theta}}), i.e.

Second, we will perform a partial or approximate matrix inversion given by:

(δ′\delta^{\prime}–approximate Jacobian-vector product) Let gi\bm{g}_{i} be a vector such that

where ϕi=Algi(θ){\bm{\phi}}_{i}=\mathcal{A}lg_{i}({\bm{\theta}}) and Algi\mathcal{A}lg_{i} is based on definition 1.

Note that gi\bm{g}_{i} in definition 2 is an approximation of the meta-gradient for task Ti\mathcal{T}_{i}. Observe that gi\bm{g}_{i} can be obtained as an approximate solution to the optimization problem:

The conjugate gradient (CG) algorithm is particularly well suited for this problem due to its excellent iteration complexity and requirement of only Hessian-vector products of the form ∇2L^i(ϕi)v\nabla^{2}\hat{\mathcal{L}}_{i}({\bm{\phi}}_{i})\bm{v}. Such hessian-vector products can be obtained cheaply without explicitly forming or storing the Hessian matrix (as we discuss in Appendix C). This CG based inversion has been successfully deployed in Hessian-free or Newton-CG methods for deep learning and trust region methods in reinforcement learning . Algorithm 1 presents the full practical algorithm. Note that these approximations to develop a practical algorithm introduce errors in the meta-gradient computation. We analyze the impact of these errors in Section 3.2 and show that they are controllable. See Appendix A for how iMAML generalizes prior gradient optimization based meta-learning algorithms.

2 Theory

In Section 3.1, we outlined a practical algorithm that makes approximations to the idealized update rule of Eq. 5. Here, we attempt to analyze the impact of these approximations, and also understand the computation and memory requirements of iMAML. We find that iMAML can match the minimax computational complexity of backpropagating through the path of the inner optimizer, but is substantially better in terms of memory usage. This work to our knowledge also provides the first non-asymptotic result that analyzes approximation error due to implicit gradients. Theorem 1 provides the computational and memory complexity for obtaining an ϵ\epsilon–approximate meta-gradient. We assume Li\mathcal{L}_{i} is smooth but do not require it to be convex. We assume that GiG_{i} in Eq. LABEL:eq:update_rule is strongly convex, which can be made possible by appropriate choice of λ\lambda. The key to our analysis is a second order Lipshitz assumption, i.e. L^i(⋅)\hat{\mathcal{L}}_{i}(\cdot) is ρ\rho-Lipshitz Hessian. This assumption and setting has received considerable attention in recent optimization and deep learning literature .

Table 1 summarizes our complexity results and compares with MAML and truncated backpropagation through the path of the inner optimizer. We use κ\kappa to denote the condition number of the inner problem induced by GiG_{i} (see Equation LABEL:eq:update_rule), which can be viewed as a measure of hardness of the inner optimization problem. Mem(∇L^i)\textrm{Mem}({\nabla\hat{\mathcal{L}}_{i}}) is the memory taken to compute a single derivative ∇L^i\nabla\hat{\mathcal{L}}_{i}. Under the assumption that Hessian vector products are computed with the reverse mode of autodifferentiation, we will have that both: the compute time and memory used for computing a Hessian vector product are with a (universal) constant factor of the compute time and memory used for computing ∇L^i\nabla\hat{\mathcal{L}}_{i} itself (see Appendix C). This allows us to measure the compute time in terms of the number of ∇L^i\nabla\hat{\mathcal{L}}_{i} computations. We refer readers to Appendix D for additional discussion about the algorithms and their trade-offs.

(Informal Statement; Approximation error in Algorithm 2) Suppose that: Li(⋅)\mathcal{L}_{i}(\cdot) is BB Lipshitz and LL smooth function; that Gi(⋅,θ)G_{i}(\cdot,{\bm{\theta}}) (in Eq. LABEL:eq:update_rule) is a μ\mu-strongly convex function with condition number κ\kappa; that DD is the diameter of search space for ϕ{\bm{\phi}} in the inner optimization problem (i.e. ∥Algi⋆(θ)∥≤D\|\mathcal{A}lg^{\star}_{i}({\bm{\theta}})\|\leq D); and L^i(⋅)\hat{\mathcal{L}}_{i}(\cdot) is ρ\rho-Lipshitz Hessian.

Let gi\bm{g}_{i} be the task meta-gradient returned by Algorithm 2. For any task ii and desired accuracy level ϵ\epsilon, Algorithm 2 computes an approximate task-specific meta-gradient with the following guarantee:

The formal statement of the theorem and the proof are provided the appendix. Importantly, the algorithm’s memory requirement is equivalent to the memory needed for Hessian-vector products which is a small constant factor over the memory required for gradient computations, assuming the reverse mode of auto-differentiation is used. Finally, the next corollary shows that iMAML efficiently finds a stationary point of F(⋅)F(\cdot), due to iMAML having controllable exact-solve error.

(iMAML finds stationary points) Suppose the conditions of Theorem 1 hold and that F(⋅)F(\cdot) is an LFL_{F} smooth function. Then the implicit MAML algorithm (Algorithm 1), when the batch size is MM (so that we are doing gradient descent), will find a point θ{\bm{\theta}} such that:

Experimental Results and Discussion

In our experimental evaluation, we aim to answer the following questions empirically: (1) Does the iMAML algorithm asymptotically compute the exact meta-gradient? (2) With finite iterations, does iMAML approximate the meta-gradient more accurately compared to MAML? (3) How does the computation and memory requirements of iMAML compare with MAML? (4) Does iMAML lead to better results in realistic meta-learning problems? We have answered (1) - (3) through our theoretical analysis, and now attempt to validate it through numerical simulations. For (1) and (2), we will use a simple synthetic example for which we can compute the exact meta-gradient and compare against it (exact-solve error, see definition 3). For (3) and (4), we will use the common few-shot image recognition domains of Omniglot and Mini-ImageNet.

To study the question of meta-gradient accuracy, Figure 2 considers a synthetic regression example, where the predictions are linear in parameters. This provides an analytical expression for Algi⋆\mathcal{A}lg^{\star}_{i} allowing us to compute the true meta-gradient. We fix gradient descent (GD) to be the inner optimizer for both MAML and iMAML. The problem is constructed so that the condition number (κ)(\kappa) is large, thereby necessitating many GD steps. We find that both iMAML and MAML asymptotically match the exact meta-gradient, but iMAML computes a better approximation in finite iterations. We observe that with 2 CG iterations, iMAML incurs a small terminal error. This is consistent with our theoretical analysis. In Algorithm 2, δ\delta is dominated by δ′\delta^{\prime} when only a small number of CG steps are used. However, the terminal error vanishes with just 5 CG steps. The computational cost of 1 CG step is comparable to 1 inner GD step with the MAML algorithm, since both require 1 hessian-vector product (see section C for discussion). Thus, the computational cost as well as memory of iMAML with 100 inner GD steps is significantly smaller than MAML with 100 GD steps.

To study (3), we turn to the Omniglot dataset which is a popular few-shot image recognition domain. Figure 2 presents compute and memory trade-off for MAML and iMAML (on 20-way, 5-shot Omniglot). Memory for iMAML is based on Hessian-vector products and is independent of the number of GD steps in the inner loop. The memory use is also independent of the number of CG iterations, since the intermediate computations need not be stored in memory. On the other hand, memory for MAML grows linearly in grad steps, reaching the capacity of a 12 GB GPU in approximately 16 steps. First-order MAML (FOMAML) does not back-propagate through the optimization process, and thus the computational cost is only that of performing gradient descent, which is needed for all the algorithms. The computational cost for iMAML is also similar to FOMAML along with a constant overhead for CG that depends on the number of CG steps. Note however, that FOMAML does not compute an accurate meta-gradient, since it ignores the Jacobian. Compared to FOMAML, the compute cost of MAML grows at a faster rate. FOMAML requires only gradient computations, while backpropagating through GD (as done in MAML) requires a Hessian-vector products at each iteration, which are more expensive.

Finally, we study empirical performance of iMAML on the Omniglot and Mini-ImageNet domains. Following the few-shot learning protocol in prior work , we run the iMAML algorithm on the dataset for different numbers of class labels and shots (in the N-way, K-shot setting), and compare two variants of iMAML with published results of the most closely related algorithms: MAML, FOMAML, and Reptile. While these methods are not state-of-the-art on this benchmark, they provide an apples-to-apples comparison for studying the use of implicit gradients in optimization-based meta-learning. For a fair comparison, we use the identical convolutional architecture as these prior works. Note however that architecture tuning can lead to better results for all algorithms .

The first variant of iMAML we consider involves solving the inner level problem (the regularized objective function in Eq. LABEL:eq:update_rule) using gradient descent. The meta-gradient is computed using conjugate gradient, and the meta-parameters are updated using Adam. This presents the most straightforward comparison with MAML, which would follow a similar procedure, but backpropagate through the path of optimization as opposed to invoking implicit differentiation. The second variant of iMAML uses a second order method for the inner level problem. In particular, we consider the Hessian-free or Newton-CG method. This method makes a local quadratic approximation to the objective function (in our case, G(ϕ′,θ)G({{\bm{\phi}}^{\prime}},{\bm{\theta}}) and approximately computes the Newton search direction using CG. Since CG requires only Hessian-vector products, this way of approximating the Newton search direction is scalable to large deep neural networks. The step size can be computed using regularization, damping, trust-region, or linesearch. We use a linesearch on the training loss in our experiments to also illustrate how our method can handle non-differentiable inner optimization loops. We refer the readers to Nocedal & Wright and Martens for a more detailed exposition of this optimization algorithm. Similar approaches have also gained prominence in reinforcement learning .

Tables 2 and 3 present the results on Omniglot and Mini-ImageNet, respectively. On the Omniglot domain, we find that the GD version of iMAML is competitive with the full MAML algorithm, and substatially better than its approximations (i.e., first-order MAML and Reptile), especially for the harder 20-way tasks. We also find that iMAML with Hessian-free optimization performs substantially better than the other methods, suggesting that powerful optimizers in the inner loop can offer benifits to meta-learning. In the Mini-ImageNet domain, we find that iMAML performs better than MAML and FOMAML. We used λ=0.5\lambda=0.5 and 1010 gradient steps in the inner loop. We did not perform an extensive hyperparameter sweep, and expect that the results can improve with better hyperparameters. 55 CG steps were used to compute the meta-gradient. The Hessian-free version also uses 55 CG steps for the search direction. Additional experimental details are Appendix F.

Related Work

Our work considers the general meta-learning problem , including few-shot learning . Meta-learning approaches can generally be categorized into metric-learning approaches that learn an embedding space where non-parametric nearest neighbors works well , black-box approaches that train a recurrent or recursive neural network to take datapoints as input and produce weight updates or predictions for new inputs , and optimization-based approaches that use bi-level optimization to embed learning procedures, such as gradient descent, into the meta-optimization problem . Hybrid approaches have also been considered to combine the benefits of different approaches . We build upon optimization-based approaches, particularly the MAML algorithm , which meta-learns an initial set of parameters such that gradient-based fine-tuning leads to good generalization. Prior work has considered a number of inner loops, ranging from a very general setting where all parameters are adapted using gradient descent , to more structured and specialized settings, such as ridge regression , Bayesian linear regression , and simulated annealing . The main difference between our work and these approaches is that we show how to analytically derive the gradient of the outer objective without differentiating through the inner learning procedure.

Mathematically, we view optimization-based meta-learning as a bi-level optimization problem. Such problems have been studied in the context of few-shot meta-learning (as discussed previously), gradient-based hyperparameter optimization , and a range of other settings . Some prior works have derived implicit gradients for related problems while others propose innovations to aid back-propagation through the optimization path for specific algorithms , or approximations like truncation . While the broad idea of implicit differentiation is well known, it has not been empirically demonstrated in the past for learning more than a few parameters (e.g., hyperparameters), or highly structured settings such as quadratic programs . In contrast, our method meta-trains deep neural networks with thousands of parameters. Closest to our setting is the recent work of Lee et al. , which uses implicit differentiation for quadratic programs in a final SVM layer. In contrast, our formulation allows for adapting the full network for generic objectives (beyond hinge-loss), thereby allowing for wider applications.

We also note that prior works involving implicit differentiation make a strong assumption of an exact solution in the inner level, thereby providing only asymptotic guarantees. In contrast, we provide finite time guarantees which allows us to analyze the case where the inner level is solved approximately. In practice, the inner level is likely to be solved using iterative optimization algorithms like gradient descent, which only return approximate solutions with finite iterations. Thus, this paper places implicit gradient methods under a strong theoretical footing for practically use.

Conclusion

In this paper, we develop a method for optimization-based meta-learning that removes the need for differentiating through the inner optimization path, allowing us to decouple the outer meta-gradient computation from the choice of inner optimization algorithm. We showed how this gives us significant gains in compute and memory efficiency, and also conceptually allows us to use a variety of inner optimization methods. While we focused on developing the foundations and theoretical analysis of this method, we believe that this work opens up a number of interesting avenues for future study.

Broader classes of inner loop procedures. While we studied different gradient-based optimization methods in the inner loop, iMAML can in principle be used with a variety of inner loop algorithms, including dynamic programming methods such as QQ-learning, two-player adversarial games such as GANs, energy-based models , and actor-critic RL methods, and higher-order model-based trajectory optimization methods. This significantly expands the kinds of problems that optimization-based meta-learning can be applied to.

Acknowledgements

Aravind Rajeswaran thanks Emo Todorov for valuable discussions about implicit gradients and potential application domains; Aravind Rajeswaran also thanks Igor Mordatch and Rahul Kidambi for helpful discussions and feedback. Sham Kakade acknowledges funding from the Washington Research Foundation for innovation in Data-intensive Discovery; Sham Kakade also graciously acknowledges support from ONR award N00014-18-1-2247, NSF Award CCF-1703574, and NSF CCF 1740551 award.

References

Appendix A Relationship between iMAML and Prior Algorithms

The presented iMAML algorithm has close connections, as well as notable differences, to a number of related algorithms like MAML , first-order MAML, and Reptile . Conventionally, these algorithms do not consider any explicit regularization in the inner-level and instead rely on early stopping, through only a few gradient descent steps. In our problem setting described in Eq. LABEL:eq:update_rule, we consider an explicitly regularized inner-level problem (refer to discussion in Section 2.2). We describe the connections between the algorithms in this explicitly regularized setting below.

MAML. The MAML algorithm first invokes an iterative algorithm to solve the inner optimization problem (see definition 1). Subsequently, it backpropagates through the path of the optimization algorithm to update the meta-parameters as:

Since Algi(θ)\mathcal{A}lg_{i}({\bm{\theta}}) approximates Algi⋆(θ)\mathcal{A}lg^{\star}_{i}({\bm{\theta}}), it can be viewed that both MAML and iMAML intend to perform the same idealized update in Eq. 5. However, they perform the meta-gradient computation very differently. MAML backpropagates through the path of an iterative algorithm, while iMAML computes the meta-gradient through the implicit Jacobian approach outlined in Section 3.1 (see Figure 1 for a visual depiction). As a result, iMAML can be vastly more efficient in memory while having lesser or comparable computational requirements. It also allows for higher order optimization methods and non-differentiable components.

First-order MAML ignores the effect of meta-parameters θ{\bm{\theta}} on task parameters {ϕi}\{{\bm{\phi}}_{i}\} in the meta-gradient computation and updates the meta-parameters as:

Note that iMAML strictly generalizes this, since first-order MAML is simply iMAML when the conjugate gradient procedure is not invoked (or corresponds to 0 steps of CG). Thus, iMAML allows for an easy way to interpolate from first-order MAML to the full MAML algorithm.

Reptile , similar to first-order MAML, ignores the dependence of task-parameters on meta-parameters. However, instead of following the gradients at ϕi=Algi(θk){\bm{\phi}}_{i}=\mathcal{A}lg_{i}({\bm{\theta}}^{k}), Reptile uses the task-parameters as targets and slowly moves meta-parameters towards them:

From the proximal point equation in the proof of Lemma 1, we have ϕi=θk−1λ∇ϕLi(ϕi){\bm{\phi}}_{i}={\bm{\theta}}^{k}-\frac{1}{\lambda}\nabla_{\bm{\phi}}\mathcal{L}_{i}({\bm{\phi}}_{i}), using which we see that the Reptile equation becomes: θk+1=θk−ηλM∑i=1M∇ϕLi(ϕi){\bm{\theta}}^{k+1}={\bm{\theta}}^{k}-\frac{\eta}{\lambda M}\sum_{i=1}^{M}\nabla_{\bm{\phi}}\mathcal{L}_{i}({\bm{\phi}}_{i}). Thus, Reptile and first-order MAML are identical in our problem formulation up to the choice of learning rate. Making the regularization explicit allows us to illustrate this equivalence.

Appendix B Optimization Preliminaries

where ∥⋅∥\|\cdot\| denotes the spectral norm.

We will make use of the following black-box complexity of first-order gradient methods for minimizing strongly convex and smooth functions.

using a number of gradient computations of ff that is bounded as follows:

Appendix C Review: Time and Space Complexity of Hessian-Vector Products

We briefly discuss the time and space complexity of Hessian-vector product computation using the reverse mode of automatic differentiation. The reverse mode of automatic differentiation is the widely used method for automatic differentiation in modern software packages like TensorFlow and PyTorch . Recall that for a differentiable function f(x)f(x), the reverse mode of automatic differentiation computes ∇f(x)\nabla f(x) in time that is no more than a factor of 55 of the time it takes to compute f(x)f(x) itself (see for review). As our algorithm makes use of Hessian vector products, we will make use of the following assumption as to how Hessian vector products will be computed when executing Algorithm 2.

(Complexity of Hessian-vector product) We assume that the time to compute the Hessian-vector product ∇ϕ2L^i(ϕ)v\nabla^{2}_{\bm{\phi}}\hat{\mathcal{L}}_{i}({\bm{\phi}})\bm{v} is no more than a (universal) constant over the time used to compute ∇L^i(ϕ)\nabla\hat{\mathcal{L}}_{i}({\bm{\phi}}) (typically, this constant is 55). Furthermore, we assume that the memory used to compute the Hessian-vector product ∇ϕ2L^i(ϕ)v\nabla^{2}_{\bm{\phi}}\hat{\mathcal{L}}_{i}({\bm{\phi}})\bm{v} is no more than twice the memory used when computing ∇L^i(ϕ)\nabla\hat{\mathcal{L}}_{i}({\bm{\phi}}). This assumption is valid if the reverse mode of automatic differentiation is used to compute Hessian vector products (see ).

A few remarks about this assumption are in order. With regards to computation, first observe that the gradient of the scalar function ∇ϕL^i(ϕ)⊤v\nabla_{\bm{\phi}}\hat{\mathcal{L}}_{i}({\bm{\phi}})^{\top}\bm{v} is the desired Hessian vector product ∇ϕ2L^i(ϕ)v\nabla_{\bm{\phi}}^{2}\hat{\mathcal{L}}_{i}({\bm{\phi}})\bm{v}. Thus computing the Hessian vector product using the reverse mode is within a constant factor of computing the function itself, which is simply the cost of computing ∇L^i(ϕ)⊤v\nabla\hat{\mathcal{L}}_{i}({\bm{\phi}})^{\top}\bm{v}. The issue of memory is more subtle (see ), which we now discuss. The memory used to compute the gradient of a scalar cost function f(x)f(x) using the reverse mode of auto-differentiation is proportional to the size of the computation graph; precisely, the memory required to compute the gradient is equal to the total space required to store all the intermediate variables used when computing f(x)f(x). In practice, this is often much larger than the memory required to compute f(x)f(x) itself, due to that all intermediate variables need not be simultaneously stored in memory when computing f(x)f(x). However, for the special case of computing the gradient of the function f(ϕ)=∇ϕL^i(ϕ)⊤vf({\bm{\phi}})=\nabla_{\bm{\phi}}\hat{\mathcal{L}}_{i}({\bm{\phi}})^{\top}\bm{v}, the factor of 22 in the memory bound is a consequence of the following reason: first, using the reverse mode to compute f(ϕ)f({\bm{\phi}}) means we already have stored the computation graph of L^i(ϕ)\hat{\mathcal{L}}_{i}({\bm{\phi}}) itself. Furthermore, the size of the computation graph for computing f(ϕ)=∇ϕL^i(ϕ)⊤vf({\bm{\phi}})=\nabla_{\bm{\phi}}\hat{\mathcal{L}}_{i}({\bm{\phi}})^{\top}\bm{v} is essentially the same size as the computation graph of L^i(ϕ)\hat{\mathcal{L}}_{i}({\bm{\phi}}). This leads to the factor of 22 memory bound; see Griewank for further discussion.

Appendix D Additional Discussion About Compute and Memory Complexity

Our main complexity results are summarized in Table 1. For these results, we consider two notions of error that are subtly different, which we explicitly define below. Let gi\bm{g}_{i} be the computed meta-gradient for task Ti\mathcal{T}_{i}. Then, the errors we consider are:

Exact-solve error (our notion of error): Our goal is to accurately compute the gradient of F(θ)F(\theta) as defined in Equation LABEL:eq:update_rule, where Algi⋆(θ)\mathcal{A}lg^{\star}_{i}(\theta) is an exact algorithm. Specifically, we seek to compute a gi\bm{g}_{i} such that:

where ϵ\epsilon is the error in the gradient computation.

Approx-solve error: Here we suppose that Algi\mathcal{A}lg_{i} computes a δ\delta–accurate solution to the inner optimization problem over GiG_{i} in Eq. LABEL:eq:update_rule, i.e. that Algi\mathcal{A}lg_{i} satisfies ∥Algi(θ)−Algi⋆(θ)∥≤δ\|\mathcal{A}lg_{i}({\bm{\theta}})-\mathcal{A}lg^{\star}_{i}({\bm{\theta}})\|\leq\delta, as per definition 1. Then the objective is to compute a g\bm{g} such that:

where ϵ\epsilon is the error in the gradient computation of dθLi(Algi(θ))\bm{d}_{\bm{\theta}}\mathcal{L}_{i}(\mathcal{A}lg_{i}({\bm{\theta}})). Subtly, note that the gradient is with respect to the δ\delta-approximate algorithm, as opposed to using Algi⋆\mathcal{A}lg^{\star}_{i}.

For the complexity results, we assume that MAML invokes Algi\mathcal{A}lg_{i} to get a δ\delta-approximate solution for inner problem (recall definition 1). The exact-solve error for MAML is not known in the literature; in particular, even as δ→0\delta\rightarrow 0 it is not evident if the approx-solve solution tends to the exact-solve solution, unless further regularity conditions are imposed. The approx-solve error for MAML is , ignoring finite-precision and numerical issues, since it backpropagates through the path. Truncated backprop also invokes Algi\mathcal{A}lg_{i} to obtain a δ\delta-approximate solution but instead performs a truncated or partial back-propagation so that it uses a smaller number of iterations when computing the gradient through the path of Algi(θ)\mathcal{A}lg_{i}({\bm{\theta}}). Exact-solve error for truncated backprop is also not known, but a small approx-solve error can be obtained with less memory than full back-prop. We use Prop 3.1 of Shaban et al. to provide a guarantee that leads to an ϵ\epsilon–accurate approximation of the full-backprop (i.e. MAML) gradient. It is not evident how accurate the truncated procedure is when an accelerated method is used instead. Finally, our iMAML algorithm also invokes an approximate solver Algi\mathcal{A}lg_{i} rather than Algi⋆\mathcal{A}lg^{\star}_{i}. However, importantly, we guarantee a small exact-solve error even though we do not require access to Algi⋆\mathcal{A}lg^{\star}_{i}. Furthermore, the iMAML algorithm also requires substantially less memory. Up to small constant factors, it only utilizes the memory required for computing a single gradient of L^i(⋅)\hat{\mathcal{L}}_{i}(\cdot).

Appendix E Proofs

Lemma 1, restated. Consider Algi⋆(θ)\mathcal{A}lg^{\star}_{i}({\bm{\theta}}) as defined in Eq. LABEL:eq:update_rule for task Ti\mathcal{T}_{i}. Let ϕi=Algi⋆(θ){\bm{\phi}}_{i}=\mathcal{A}lg^{\star}_{i}({\bm{\theta}}) be the result of Algi⋆(θ)\mathcal{A}lg^{\star}_{i}({\bm{\theta}}). If (I+1λ∇ϕ2L^i(ϕi))\left(\bm{I}+\frac{1}{\lambda}\nabla_{\bm{\phi}}^{2}\hat{\mathcal{L}}_{i}({\bm{\phi}}_{i})\right) is invertible, then the derivative Jacobian is

which is an implicit equation that often arises in proximal point methods. When the derivative exists, we can differentiate the above equation to obtain:

(Regularity conditions) Suppose the following holds for all tasks ii:

Li(⋅)\mathcal{L}_{i}(\cdot) is BB Lipshitz and LL smooth.

For all θ{\bm{\theta}}, Gi(⋅,θ)G_{i}(\cdot,{\bm{\theta}}) is both a β\beta-smooth function and a μ\mu-strongly convex function. Define:

L^i(⋅)\hat{\mathcal{L}}_{i}(\cdot) is ρ\rho-Lipshitz Hessian, i.e. ∇2L^i(⋅)\nabla^{2}\hat{\mathcal{L}}_{i}(\cdot) is ρ\rho-Lipshitz.

For all θ{\bm{\theta}}, suppose the arg-minimizer of Gi(⋅,θ)G_{i}(\cdot,{\bm{\theta}}) is unique and bounded in a ball of radius DD, i.e. for all θ{\bm{\theta}},

(Implicit Gradient Accuracy) Suppose Assumption 2 holds. Fix a task ii. Suppose that ϕi{\bm{\phi}}_{i} satisfies:

Assuming that δ<μ/(2ρ)\delta<\mu/(2\rho), we have that:

where the first inequality uses the triangle inequality.

We now bound each of these terms. For the second term,

where we the second inequality uses that ∇ϕL\nabla_{\bm{\phi}}\mathcal{L} is LL-smooth and the final inequality uses that GG is μ\mu strongly convex.

using that ∇ϕL\nabla_{\bm{\phi}}\mathcal{L} is BB Lipshitz. Now let

Due to that ∇2L^(⋅)\nabla^{2}\hat{\mathcal{L}}(\cdot) is Lipshitz Hessian, ∥Δ∥≤ρδ\|\Delta\|\leq\rho\delta. Also, by our assumption on δ\delta, we have that:

which implies that ∥(I+M−1Δ)−1∥≤2\|\left(\bm{I}+M^{-1}\Delta\right)^{-1}\|\leq 2. Hence,

The proof is completed by substitution. ∎

(Approximate Implicit Gradient Computation) Suppose Assumption 2 holds. Fix a task ii. Let

Suppose Nesterov’s accelerated gradient descent algorithm is used to compute ϕ{\bm{\phi}} (as desired in Algorithm 2), using a number of iterations that is:

and suppose Nesterov’s accelerated gradient descent algorithm (or the conjugate gradient algorithm The conjugate gradient descent algorithm also suffices and give a slightly improved iteration complexity in terms of log factors.) is used to compute gi\bm{g}_{i} using a number of iterations that is:

The result will follow from the guarantees in Lemma 2. Specifically, let us set δ=min⁡{ϵ/(2B1),μ/(2ρ)}\delta=\min\{\epsilon/(2B_{1}),\mu/(2\rho)\} and δ′=ϵ/2\delta^{\prime}=\epsilon/2. To ensure the bound of δ\delta, by Lemma 3, it suffices to use a number of iterations that is bounded by:

To ensure the bound of δ′\delta^{\prime}, the algorithm will be solving the sub-problem in Equation 7. First observe that in the context of in Lemma 2, note that ∥x⋆∥=∥(I+1λ ∇2L^i(ϕ))−1∇Li(ϕ)∥≤(λ/μ)B\|x^{\star}\|=\|\left(\bm{I}+\frac{1}{\lambda}~{}\nabla^{2}\hat{\mathcal{L}}_{i}({\bm{\phi}})\right)^{-1}\nabla\mathcal{L}_{i}({\bm{\phi}})\|\leq(\lambda/\mu)B, and so it suffices to use a number of iterations that is bounded by:

Appendix F Experiment Details

Here, we provide additional details of the experimental set-up for the experiments in Section 4. All training runs were conducted on a single NVIDIA (Titan Xp) GPU.

For the synthetic experiments, we consider a linear regression problem. We consider parametric models of the form hϕ(x)=ϕTxh_{\bm{\phi}}(\mathbf{x})={\bm{\phi}}^{T}\mathbf{x}, where x\mathbf{x} can either be the raw inputs or features (e.g. Fourier features) of the input. For task Ti\mathcal{T}_{i}, we can equivalently write a quadratic objective that represents the task loss as:

Thus, the exact meta-gradient can be written as

F.2 Omniglot and Mini-ImageNet experiments

We follow the standard training and evaluation protocol as in prior works .

The GD version of iMAML uses 16 gradient steps for 5-way 1-shot and 5-way 5-shot settings, and 25 gradient steps for 20-way 1-shot and 20-way 5-shot settings. A regularization strength of λ=2.0\lambda=2.0 was used for both. 55 steps of conjugate gradient was used to compute the meta-gradient for each task in the mini-batch, and the meta-gradients were averaged before taking a step with the default parameters of Adam in the outer loop.

The Hessian-free version of MAML proceeds by using Hessian-free or Newton-CG method for solving the inner optimization problem (with respect to ϕ{\bm{\phi}}) with objective Gi(ϕ,θ)G_{i}({\bm{\phi}},{\bm{\theta}}). This method proceeds by constructing a local quadratic approximation to the objective and approximately computing the Newton direction with conjugate gradient. 55 CG steps are used for this process in our experiments. This allows us to compute the search direction, following which a step size has to be picked. We pick the step size through line-search. This procedure of computing the approximate Newton direction and linesearch is repeated 33 times in our experiments to solve the inner optimization problem well.

Mini-ImageNet

For the GD version of iMAML, 10 GD steps were used with regularization strength of λ=0.5\lambda=0.5. Again, 5 CG steps are used to compute the meta-gradient. Similarly, in the Hessian-Free variant, we again use 55 CG steps to compute the search direction followed by line search. This process is repeated 33 times to solve the inner level optimization. Again, to compute the meta-gradient, 5 steps of CG are used.