Truncated Back-propagation for Bilevel Optimization

Amirreza Shaban, Ching-An Cheng, Nathan Hatch, Byron Boots

INTRODUCTION

Bilevel optimization has been recently revisited as a theoretical framework for designing and analyzing algorithms for hyperparameter optimization and meta learning . Mathematically, these problems can be formulated as a stochastic optimization problem with an equality constraint (see Section 1.1):

where ww and λ\lambda are the parameter and the hyperparameter, FF and fSf_{S} are the expected and the sampled upper-level objective, gSg_{S} is the sampled lower-level objective, and SS is a random variable called the context. The notation ≈λ\approx_{\lambda} means that w^S∗(λ)\hat{w}_{S}^{*}(\lambda) equals the unique return value of a prespecified iterative algorithm (e.g. gradient descent) that approximately finds a local minimum of gSg_{S}. This algorithm is part of the problem definition and can also be parametrized by λ\lambda (e.g. step size). The motivation to explicitly consider the approximate solution w^S∗(λ)\hat{w}_{S}^{*}(\lambda) rather than an exact minimizer wS∗w_{S}^{*} of gSg_{S} is that wS∗w_{S}^{*} is usually not available in closed form. This setup enables λ\lambda to account for the imperfections of the lower-level optimization algorithm.

Solving the bilevel optimization problem in (1) is challenging due to the complicated dependency of the upper-level problem on λ\lambda induced by w^S∗(λ)\hat{w}_{S}^{*}(\lambda). This difficulty is further aggravated when λ\lambda and ww are high-dimensional, precluding the use of black-box optimization techniques such as grid/random search and Bayesian optimization .

Recently, first-order bilevel optimization techniques have been revisited to solve these problems. These methods rely on an estimate of the Jacobian ∇λw^S∗(λ)\nabla_{\lambda}\hat{w}_{S}^{*}(\lambda) to optimize λ\lambda. Pedregosa and Gould et al. assume that w^S∗(λ)=wS∗\hat{w}_{S}^{*}(\lambda)=w_{S}^{*} and compute ∇λw^S∗(λ)\nabla_{\lambda}\hat{w}_{S}^{*}(\lambda) by implicit differentiation. By contrast, Maclaurin et al. and Franceschi et al. treat the iterative optimization algorithm in the lower-level problem as a dynamical system, and compute ∇λw^S∗(λ)\nabla_{\lambda}\hat{w}_{S}^{*}(\lambda) by automatic differentiation through the dynamical system. In comparison, the latter approach is less sensitive to the optimality of w^S∗(λ)\hat{w}_{S}^{*}(\lambda) and can also learn hyperparameters that control the lower-level optimization process (e.g. step size). However, due to superlinear time or space complexity (see Section 2.2), neither of these methods is applicable when both λ\lambda and ww are high-dimensional .

Few-step reverse-mode automatic differentiation and few-step forward-mode automatic differentiation have recently been proposed as heuristics to address this issue. By ignoring long-term dependencies, the time and space complexities to compute approximate gradients can be greatly reduced. While exciting empirical results have been reported, the theoretical properties of these methods remain unclear.

In this paper, we study the theoretical properties of these truncated back-propagation approaches. We show that, when the lower-level problem is locally strongly convex around w^S∗(λ)\hat{w}_{S}^{*}(\lambda), on-average convergence to an ϵ\epsilon-approximate stationary point is guaranteed by O(log⁡1/ϵ)O(\log 1/\epsilon)-step truncated back-propagation. We also identify additional problem structures for which asymptotic convergence to an exact stationary point is guaranteed. Empirically, we verify the utility of this strategy for hyperparameter optimization and meta learning tasks. We find that, compared to optimization with full back-propagation, optimization with truncated back-propagation usually shows competitive performance while requiring half as much computation time and significantly less memory.

The goal of hyperparameter optimization is to find hyperparameters λ\lambda for an optimization problem PP such that the approximate solution w^∗(λ)\hat{w}^{*}(\lambda) of PP has low cost c(w^∗(λ))c(\hat{w}^{*}(\lambda)) for some cost function cc. In general, λ\lambda can parametrize both the objective of PP and the algorithm used to solve PP. This setup is a special case of the bilevel optimization problem (1) where the upper-level objective cc does not depend directly on λ\lambda. In contrast to meta learning (discussed below), cc can be deterministic . See Section 4.2 for examples.

Many low-dimensional problems, such as choosing the learning rate and regularization constant for training neural networks, can be effectively solved with grid search. However, problems with thousands of hyperparameters are increasingly common, for which gradient-based methods are more appropriate .

Another important application of bilevel optimization, meta learning (or learning-to-learn) uses statistical learning to optimize an algorithm Aλ\mathcal{A}_{\lambda} over a distribution of tasks T\mathcal{T} and contexts SS:

It treats Aλ\mathcal{A}_{\lambda} as a parametric function, with hyperparameter λ\lambda, that takes task-specific context information SS as input and outputs a decision Aλ(S)\mathcal{A}_{\lambda}(S). The goal of meta learning is to optimize the algorithm’s performance cTc_{\mathcal{T}} (e.g. the generalization error) across tasks T\mathcal{T} through empirical observations. This general setup subsumes multiple problems commonly encountered in the machine learning literature, such as multi-task learning and few-shot learning .

Bilevel optimization emerges from meta learning when the algorithm computes Aλ(S)\mathcal{A}_{\lambda}(S) by internally solving a lower-level minimization problem with variable ww. The motivation to use this class of algorithms is that the lower-level problem can be designed so that, even for tasks T\mathcal{T} distant from the training set, Aλ\mathcal{A}_{\lambda} falls back upon a sensible optimization-based approach . By contrast, treating Aλ\mathcal{A}_{\lambda} as a general function approximator relies on the availability of a large amount of meta training data .

In other words, the decision is Aλ(S)=(w^S∗(λ),λ)\mathcal{A}_{\lambda}(S)=(\hat{w}_{S}^{*}(\lambda),\lambda) where w^S∗(λ)\hat{w}_{S}^{*}(\lambda) is an approximate minimizer of some function gS(w,λ)g_{S}(w,\lambda). Therefore, we can identify

BILEVEL OPTIMIZATION

2 Computing the hypergradient

Like , we treat the iterative optimization algorithm that solves the lower-level problem as a dynamical system. Given an initial condition w0=Ξ0(λ)w_{0}=\Xi_{0}(\lambda) at t=0t=0, the update rule can be written asFor notational simplicity, we consider the case where wtw_{t} is the state of (4); our derivation can be easily generalized to include other internal states, e.g. momentum.

in which Ξt\Xi_{t} defines the transition and and TT is the number iterations performed. For example, in gradient descent, Ξt+1(wt,λ)=wt−γt(λ)∇wg(wt,λ)\Xi_{t+1}(w_{t},\lambda)=w_{t}-\gamma_{t}(\lambda)\nabla_{w}g(w_{t},\lambda), where γt(λ)\gamma_{t}(\lambda) is the step size.

By unrolling the iterative update scheme (4) as a computational graph, we can view w^∗\hat{w}^{*} as a function of λ\lambda and compute the required derivative dλf\text{d}_{\lambda}f . Specifically, it can be shown by the chain ruleNote that this assumes gg is twice differentiable.

where At+1=∇wtΞt+1(wt,λ)A_{t+1}=\nabla_{w_{t}}\Xi_{t+1}(w_{t},\lambda), Bt+1=∇λΞt+1(wt,λ)B_{t+1}=\nabla_{\lambda}\Xi_{t+1}(w_{t},\lambda) for t≥0t\geq 0, and B0=dλΞ0(λ)B_{0}=\text{d}_{\lambda}\Xi_{0}(\lambda).

The computation of (5) can be implemented either in reverse mode or forward mode . Reverse-mode differentiation (RMD) computes (5) by back-propagation:

and finally dλf=h−1\text{d}_{\lambda}f=h_{-1}. Forward-mode differentiation (FMD) computes (5) by forward propagation:

TRUNCATED BACK-PROPAGATION

In this paper, we investigate approximating (5) with partial sums, which was previously proposed as a heuristic for bilevel optimization ( Eq. 3, Eq. 2). Formally, we perform KK-step truncated back-propagation (KK-RMD) and use the intermediate variable hT−Kh_{T-K} to construct an approximate gradient:

This approach requires storing only the last KK iterates wtw_{t}, and it also saves computation time. Note that KK-RMD can be combined with checkpointing for further savings, although we do not investigate this.

We first establish some intuitions about why using KK-RMD to optimize λ\lambda is reasonable. While building up an approximate gradient by truncating back-propagation in general optimization problems can lead to large bias, the bilevel optimization problem in (1) has some nice structure. Here we show that if the lower-level objective gg is locally strongly convex around w^∗\hat{w}^{*}, then the bias of hT−Kh_{T-K} can be exponentially small in KK. That is, choosing a small KK would suffice to give a good gradient approximation in finite precision. The proof is given in Appendix A.

Assume gg is β\beta-smooth, twice differentiable, and locally α\alpha-strongly convex in ww around {wT−K−1,…,wT}\{w_{T-K-1},\dots,w_{T}\}. Let Ξt+1(wt,λ)=wt−γ∇wg(wt,λ)\Xi_{t+1}(w_{t},\lambda)=w_{t}-\gamma\nabla_{w}g(w_{t},\lambda). For γ≤1β\gamma\leq\frac{1}{\beta}, it holds

where MB=max⁡t∈{0,…,T−K}∥Bt∥M_{B}=\max_{t\in\{0,\dots,T-K\}}\|B_{t}\|. In particular, if gg is globally α\alpha-strongly convex, then

Note 0≤(1−γα)<10\leq(1-\gamma\alpha)<1 since γ≤1β≤1α\gamma\leq\frac{1}{\beta}\leq\frac{1}{\alpha}. Therefore, Proposition 3.1 says that if w^∗\hat{w}^{*} converges to the neighborhood of a strict local minimum of the lower-level optimization, then the bias of using the approximate gradient of KK-RMD decays exponentially in KK. This exponentially decaying property is the main reason why using hT−Kh_{T-K} to update the hyperparameter λ\lambda works.

Next we show that, when the lower-level problem gg is second-order continuously differentiable, −hT−K-h_{T-K} actually is a sufficient descent direction. This is a much stronger property than the small bias shown in Proposition 3.1, and it is critical in order to prove convergence to exact stationary points (cf. Theorem 3.4). To build intuition, here we consider a simpler problem where gg is globally strongly convex and ∇λf=0\nabla_{\lambda}f=0. These assumptions will be relaxed in the next subsection.

Let gg be globally strongly convex and ∇λf=0\nabla_{\lambda}f=0. Assume gg is second-order continuously differentiable and BtB_{t} has full column rank for all tt. Let Ξt+1(wt,λ)=wt−γ∇wg(wt,λ)\Xi_{t+1}(w_{t},\lambda)=w_{t}-\gamma\nabla_{w}g(w_{t},\lambda). For all K≥1K\geq 1, with TT large enough and γ\gamma small enough, there exists c>0c>0, s.t. hT−K⊤dλf≥c∥∇w^∗f∥2h_{T-K}^{\top}\text{d}_{\lambda}f\geq c\|\nabla_{\hat{w}^{*}}f\|^{2}. This implies hT−Kh_{T-K} is a sufficient descent direction, i.e. hT−K⊤dλf≥Ω(∥dλf∥2)h_{T-K}^{\top}\text{d}_{\lambda}f\geq\Omega(\|\text{d}_{\lambda}f\|^{2}).

The full proof of this non-trivial result is given in Appendix B. Here we provide some ideas about why it is true. First, by Proposition 3.1, we know the bias decays exponentially. However, this alone is not sufficient to show that −hT−K-h_{T-K} is a sufficient descent direction. To show the desired result, Lemma 3.2 relies on the assumption that gg is second-order continuously differentiable and the fact that using gradient descent to optimize a well-conditioned function has linear convergence . These two new structural properties further reduce the bias in Proposition 3.1 and lead to Lemma 3.2. Here the full rank assumption for BtB_{t} is made to simplify the proof. We conjecture that this condition can be relaxed when K>1K>1. We leave this to future work.

2 Convergence

With these insights, we analyze the convergence of bilevel optimization with truncated back-propagation. Using Proposition 3.1, we can immediately deduce that optimizing λ\lambda with hT−Kh_{T-K} converges on-average to an ϵ\epsilon-approximate stationary point. Let ∇F(λτ)\nabla F(\lambda_{\tau}) denote the hypergradient in the τ\tauth iteration.

Suppose FF is smooth and bounded below, and suppose there is ϵ<∞\epsilon<\infty such that ∥hT−K−dλf∥≤ϵ\|h_{T-K}-\text{d}_{\lambda}f\|\leq\epsilon. Using hT−Kh_{T-K} as a stochastic first-order oracle with a decaying step size ητ=O(1/τ)\eta_{\tau}=O(1/\sqrt{\tau}) to update λ\lambda with gradient descent, it follows after RR iterations,

That is, under the assumptions in Proposition 3.1, learning with hT−Kh_{T-K} converges to an ϵ\epsilon-approximate stationary point, where ϵ=O((1−γα)−K)\epsilon=O((1-\gamma\alpha)^{-K}).

We see that the bias becomes small as KK increases. As a result, it is sufficient to perform KK-step truncated back-propagation with K=O(log⁡1/ϵ)K=O(\log 1/\epsilon) to update λ\lambda.

Next, using Lemma 3.2, we show that the bias term in Theorem 3.3 can be removed if the problem is more structured. As promised, we relax the simplifications made in Lemma 3.2 into assumptions 2 and 3 below and only assume gg is locally strongly convex.

Under the assumptions in Proposition 3.1 and Theorem 3.3, if in addition

gg is second-order continuously differentiable

BtB_{t} has full column rank around wTw_{T}

∇λf⊤(dλf+hT−K−∇λf)≥Ω(∥∇λf∥2)\nabla_{\lambda}f^{\top}(\text{d}_{\lambda}f+h_{T-K}-\nabla_{\lambda}f)\geq\Omega(\|\nabla_{\lambda}f\|^{2})

the problem is deterministic (i.e. F=fF=f)

then for all K≥1K\geq 1, with TT large enough and γ\gamma small enough, the limit point is an exact stationary point, i.e. lim⁡τ→∞∥∇F(λτ)∥=0\lim_{\tau\to\infty}\|\nabla F(\lambda_{\tau})\|=0.

Theorem 3.4 shows that if the partial derivative ∇λf\nabla_{\lambda}f does not interfere strongly with the partial derivative computed through back-propagating the lower-level optimization procedure (assumption 3), then optimizing λ\lambda with hT−Kh_{T-K} converges to an exact stationary point. This is a very strong result for an interesting special case. It shows that even with one-step back-propagation hT−1h_{T-1}, updating λ\lambda can converge to a stationary point.

This non-interference assumption unfortunately is necessary; otherwise, truncating the full RMD leads to constant bias, as we show below (proved in Appendix E).

There is a problem, satisfying all but assumption 3 in Theorem 3.4, such that optimizing λ\lambda with hT−Kh_{T-K} does not converge to a stationary point.

Note however that the non-interference assumption is satisfied when ∇λf=0\nabla_{\lambda}f=0, i.e. when the upper-level problem does not directly depend on the hyperparameter. This is the case for many practical applications: e.g. hyperparameter optimization, meta-learning regularization models, image desnosing , data hyper-cleaning , and task interaction .

3 Relationship with implicit differentiation

The gradient estimate hT−Kh_{T-K} is related to implicit differentiation, which is a classical first-order approach to solving bilevel optimization problems . Assume gg is second-order continuously differentiable and that its optimal solution uniquely exists such that w∗=w∗(λ)w^{*}=w^{*}(\lambda). By the implicit function theorem , the total derivative of ff with respect to λ\lambda can be written as

Here we show that, in the limit where w^∗\hat{w}^{*} converges to w∗w^{*}, hT−Kh_{T-K} can be viewed as approximating the matrix inverse in (11) with an order-KK Taylor series. This can be seen from the next proposition.

Under the assumptions in Proposition 3.1, suppose wtw_{t} converges to a stationary point w∗w^{*}. Let A∞=lim⁡t→∞AtA_{\infty}=\lim_{t\to\infty}A_{t} and B∞=lim⁡t→∞BtB_{\infty}=\lim_{t\to\infty}B_{t}. For γ<1β\gamma<\frac{1}{\beta}, it satisfies that

By Proposition 3.6, we can write dλf\text{d}_{\lambda}f in (11) as

That is, hT−Kh_{T-K} captures the first KK terms in the Taylor series, and the residue term has an upper bound as in Proposition 3.1.

Given this connection, we can compare the use of hT−Kh_{T-K} and approximating (11) using KK steps of conjugate gradient descent for high-dimensional problems . First, both approaches require local strong-convexity to ensure a good approximation. Specifically, let κ=βα>0\kappa=\frac{\beta}{\alpha}>0 locally around the limit. Using hT−Kh_{T-K} has a bias in O((1−1κ)K)O((1-\frac{1}{\kappa})^{K}), whereas using (11) and inverting the matrix with KK iterations of conjugate gradient has a bias in O((1−1κ)K)O((1-\frac{1}{\sqrt{\kappa}})^{K}) . Therefore, when w∗w^{*} is available, solving (11) with conjugate gradient descent is preferable. However, in practice, this is hardly true. When an approximate solution w^∗\hat{w}^{*} to the lower-level problem is used, adopting (11) has no control on the approximate error, nor does it necessarily yield a descent direction. On the contrary, hT−Kh_{T-K} is based on Proposition 3.1, which uses a weaker assumption and does not require the convergence of wtw_{t} to a stationary point. Truncated back-propagation can also optimize the hyperparameters that control the lower-level optimization process, which the implicit differentiation approach cannot do.

EXPERIMENTS

This deterministic problem satisfies all of the assumptions in the previous section, particularly those of Theorem 3.4: gg is 11-smooth and 12\frac{1}{2}-strongly convex, with

and B0=0B_{0}=0. Although ff is somewhat complicated, with many saddle points, it satisfies the non-interference assumption because ∇λf=0\nabla_{\lambda}f=0.

Figure 1 visualizes Proposition 3.1 by plotting the approximation error ∥hT−K−dλf∥\|h_{T-K}-\text{d}_{\lambda}f\| and the theoretical bound (1−γα)Kγα∥∇w^∗f∥MB\frac{(1-\gamma\alpha)^{K}}{\gamma\alpha}\|\nabla_{\hat{w}^{*}}f\|M_{B} at λ=(1,1)\lambda=(1,1). For this problem, α=12\alpha=\frac{1}{2}, MB=∥γG∥=γM_{B}=\|\gamma G\|=\gamma, and ∇w^∗f\nabla_{\hat{w}^{*}}f can be found analytically from w^∗=Cw0+(I−C)λ\hat{w}^{*}=Cw_{0}+(I-C)\lambda, where C=(I−γG)TC=(I-\gamma G)^{T}. Figure 4 (left) plots the iterates λτ\lambda_{\tau} when optimizing ff using 11-RMD and a decaying meta-learning rate ητ=η0τ\eta_{\tau}=\frac{\eta_{0}}{\sqrt{\tau}}. Because ∥hT−K∥\|h_{T-K}\| varies widely with KK, we tune η0\eta_{0} to ensure that the first update η1hT−K(λ1)\eta_{1}h_{T-K}(\lambda_{1}) has norm 0.60.6. In comparison with the true gradient dλf\text{d}_{\lambda}f at these points, we see that hT−1h_{T-1} is indeed a descent direction. Figure 2 (left) visualizes this in a different way, by plotting hT−K⊤dλf/∥dλf∥2h_{T-K}^{\top}\text{d}_{\lambda}f/\|\text{d}_{\lambda}f\|^{2} for various KK at each point λτ\lambda_{\tau} along the K=1K=1 trajectory. By Lemma 3.2, this ratio stays well away from zero.

To demonstrate the biased convergence of Theorem 3.3, we break assumption 3 of Theorem 3.4 by changing the upper objective to f~(w^∗,λ):=f(w^∗,λ)+5∥λ−(1,0)∥2\widetilde{f}(\hat{w}^{*},\lambda):=f(\hat{w}^{*},\lambda)+5\|\lambda-(1,0)\|^{2} so that ∇λf~≠0\nabla_{\lambda}\widetilde{f}\neq 0. The guarantee of Lemma 3.2 no longer applies, and we see in Figure 2 (right) that hT−K⊤dλf/∥dλf∥2h_{T-K}^{\top}\text{d}_{\lambda}f/\|\text{d}_{\lambda}f\|^{2} can become negative. Indeed, Figure 3 shows that optimizing f~\widetilde{f} with hT−1h_{T-1} converges to a suboptimal point. However, it also shows that using larger KK rapidly decreases the bias.

For the original objective ff, Theorem 3.4 guarantees exact convergence. Figure 4 shows optimization trajectories for various KK, and a log-scale plot of their convergence rates. Note that, because the lower-level problem cannot be perfectly solved within TT steps, the optimal λ\lambda is offset from the origin. Truncated back-propagation can handle this, but it breaks the assumptions required by the implicit differentiation approach to bilevel optimization.

2 Hyperparameter optimization problems

We optimize the lower-level problem gg through T=100T=100 steps of gradient descent with γ=1\gamma=1 and consider how adjusting KK changes the performance of KK-RMD. See Appendix G.1 for more experimental setup. Our hypothesis is that KK-RMD for small KK works almost as well as full RMD in terms of validation and test accuracy, while requiring less time and far less memory. We also hypothesize that KK-RMD does almost as well as full RMD in identifying which samples were corrupted . Because our formulation of the problem is unconstrained, the weights σ(λi)\sigma(\lambda_{i}) are never exactly zero. However, we can calculate an F1 score by setting a threshold on λ\lambda: if σ(λi)<σ(−3)≈0.047\sigma(\lambda_{i})<\sigma(-3)\approx 0.047, then the hyper-cleaner has marked example ii as corrupted.F1 scores for other choices of the threshold were very similar. See Appendix G.1 for details.

Table 2 reports these metrics for various KK. We see that 11-RMD is somewhat worse than the others, and that validation loss (the outer objective ff) decreases with KK more quickly than generalization error. The F1 score is already maximized at K=5K=5. These preliminary results indicate that in situations with limited memory, KK-RMD for small KK (e.g. K=5K=5) may be a reasonable fallback: it achieves results close to full backprop, and it runs about twice as fast.

From a theoretical optimization perspective, we wonder whether KK-RMD converges to a stationary point of ff. Data hypercleaning satisfies all of the assumptions of Theorem 3.4 except that BtB_{t} is not full column rank (since M<NM<N). In particular, the validation loss ff is deterministic and satisfies ∇λf=0\nabla_{\lambda}f=0. Figure 5 plots the norm of the true gradient dλf\text{d}_{\lambda}f on a log scale at the KK-RMD iterates for various KK. We see that, despite satisfying almost all assumptions, this problem exhibits biased convergence. The limit of ∥dλf∥\|\text{d}_{\lambda}f\| decreases slowly with KK, but recall from Table 2 that practical metrics improve more quickly.

2.2 Task interaction

We next consider the problem of multitask learning . Similar to , we formulate this as a hyperparameter optimization problem as follows. The lower-level objective g(w,{C,ρ})g(w,\{C,\rho\}) learns VV different linear models with parameter set w={wv}v=1Vw=\{w_{v}\}_{v=1}^{V}:

where l(w)l(w) is the training loss of the multi-class linear logistic regression model, ρ\rho is a regularization constant, and CC is a nonnegative, symmetric hyperparameter matrix that encodes the similarity between each pair of tasks. After 100100 iterations of gradient descent with learning rate 0.10.1, this yields w^∗\hat{w}^{*}. The upper-level objective c(w^∗)c(\hat{w}^{*}) estimates the linear regression loss of the learned model w^∗\hat{w}^{*} on a validation set. Presumably, this will be improved by tuning CC to reflect the true similarities between the tasks. The tasks that we consider are image recognition trained on very small subsets of the datasets CIFAR-1010 and CIFAR-100100. See Appendix G.2 for more details.

From an optimization standpoint, we are most interested in the upper-level loss on the validation set, since that is what is directly optimized, and its value is a good indication of the performance of the inexact gradient. Figure 6 plots this learning curve along with two other metrics of theoretical interest: norm of the true gradient, and cosine similarity between the true and approximate gradients. In CIFAR100, the validation error and gradient norm plots show that KK-RMD converges to an approximate stationary point with a bias that rapidly decreases as KK increases, agreeing with Proposition 3.1. Also, we find that negative values exist in the cosine similarity of 11-RMD, which implies that not all the assumptions in Theorem 3.4 hold for this problem (e.g. BtB_{t} might not be full rank, or the the inner problem might not be locally strong convex around w^∗\hat{w}^{*}.) In CIFAR10, some unusal behavior happens. For K>1K>1, the truncated gradient and the full gradient directions eventually become almost the same. We believe this is a very interesting observation but beyond the scope of the paper to explain.

In Table 3, we report the testing accuracy over 10 trials. While in general increasing the number of back-propagation steps improves accuracy, the gaps are small. A thorough investigation of the relationship between convergence and generalization is an interesting open question of both theoretical and practical importance.

3 Meta-learning: One-shot classification

The aim of this experiment is to evaluate the performance of truncated back-propagation in multi-task, stochastic optimization problems. We consider in particular the one-shot classification problem , where each task T\mathcal{T} is a kk-way classification problem and the goal is learn a hyperparameter λ\lambda such that each task can be solved with few training samples.

In each hyper-iteration, we sample a task, a training set, and a validation set as follows: First, kk classes are randomly chosen from a pool of classes to define the sampled task T\mathcal{T}. Then the training set S={(xi,yi)}i=1kS=\{(x_{i},y_{i})\}_{i=1}^{k} is created by randomly drawing one training example (xi,yi)(x_{i},y_{i}) from each of the kk classes. The validation set QQ is constructed similarly, but with more examples from each class. The lower-level objective gS(w,λ)g_{S}(w,\lambda) is

where l(⋅,⋅)l(\cdot,\cdot) is the kk-way cross-entropy loss, and nn(⋅;w,λ)nn(\cdot;w,\lambda) is a deep neural network parametrized by w={w1,…,wV}w=\{w_{1},\dots,w_{V}\} and optionally hyperparameter λ\lambda. To prevent overfitting in the lower-level optimization, we regularize each parameter wjw_{j} to be close to center cjc_{j} with weight ρj>0\rho_{j}>0. Both cjc_{j} and ρj\rho_{j} are hyperparameters, as well as the inner learning rate γ\gamma. The upper-level objective is the loss of the trained network on the sampled validation set QQ. In contrast to other experiments, this is a stochastic optimization problem. Also, Aλ(S)(xi)=nn(xi;w^∗,λ)\mathcal{A}_{\lambda}(S)(x_{i})=nn(x_{i};\hat{w}^{*},\lambda) depends directly on the hyperparameter λ\lambda, in addition to the indirect dependence through w^∗\hat{w}^{*} (i.e. ∇λf≠0\nabla_{\lambda}f\neq 0).

We use the Omniglot dataset and a similar neural network as used in with small modifications. Please refer to Appendix G.3 for more details about the model and the data splits. We set T=50T=50 and optimize over the hyperparameter λ={λl1,λl2,c,ρ,γ}\lambda=\{\lambda_{l_{1}},\lambda_{l_{2}},c,\rho,\gamma\}. The average accuracy of each model is evaluated over 120120 randomly sampled training and validation sets from the meta-testing dataset. For comparison, we also try using full RMD with a very short horizon T=1T=1, which is common in recent work on few-shot learning .

The statistics are shown in Table 4 and the learning curves in Figure 7. In addition to saving memory, all truncated methods are faster than full RMD, sometimes even five times faster. These results suggest that running few-step back-propagation with more hyper-iterations can be more efficient than the full RMD. To support this hypothesis, we also ran 11-RMD and 1010-RMD for an especially large number of hyper-iterations (1515k). Even with this many hyper-iterations, the total runtime is less than full RMD with 50005000 iterations, and the results are significantly improved. We also find that while using a short horizon (T=1T=1) is faster, it achieves a lower accuracy at the same number of iterations.

CONCLUSION

We analyze KK-RMD, a first-order heuristic for solving bilevel optimization problems when the lower-level optimization is itself approximated in an iterative way. We show that KK-RMD is a valid alternative to full RMD from both theoretical and empirical standpoints. Theoretically, we identify sufficient conditions for which the hyperparameters converge to an approximate or exact stationary point of the upper-level objective. The key observation is that when w^∗\hat{w}^{*} is near a strict local minimum of the lower-level objective, gradient approximation error decays exponentially with reverse depth. Empirically, we explore the properties of this optimization method with four proof-of-concept experiments. We find that although exact convergence appears to be uncommon in practice, the performance of KK-RMD is close to full RMD in terms of application-specific metrics (such as generalization error). It is also roughly twice as fast. These results suggest that in hyperparameter optimization or meta learning applications with memory constraints, truncated back-propagation is a reasonable choice.

Our experiments use a modest number of parameters MM, hyperparameters NN, and horizon length TT. This is because we need to be able to calculate both KK-RMD and full RMD in order to compare their performance. One promising direction for future research is to use KK-RMD for bilevel optimization problems that require powerful function approximators at both levels of optimization. Truncated RMD makes this approach feasible and enables comparing bilevel optimization to other meta-learning methods on difficult benchmarks.

Appendix

Appendix A Proof of Proposition 3.1

Let dλf−hT−K=eK\text{d}_{\lambda}f-h_{T-K}=e_{K}. By definition of hT−Kh_{T-K},

Therefore, when gg is locally α\alpha-strongly convex with respect to ww in the neighborhood of {wT−K−1,…,wT}\{w_{T-K-1},\dots,w_{T}\},

Suppose gg is β\beta-smooth but nonconvex. In the worst case, if the smallest eigenvalue of ∇w,wg(wt−1,λ)\nabla_{w,w}g(w_{t-1},\lambda) is −β-\beta, then ∥At∥=1+γβ≤2\|A_{t}\|=1+\gamma\beta\leq 2 for t=0,…,T−Kt=0,\dots,T-K. This gives the bound in (9). However, if gg is globally strongly convex, then

The bound (10) uses the fact that ∑t=0T−K(1−γα)t≤∑t=0∞(1−γα)t=1γα\sum_{t=0}^{T-K}(1-\gamma\alpha)^{t}\leq\sum_{t=0}^{\infty}(1-\gamma\alpha)^{t}=\frac{1}{\gamma\alpha} ∎

Appendix B Proof of Lemma 3.2

To illustrate the idea, here we prove the case where K=1K=1. For K>1K>1, similar steps can be applied. To prove the statement, we first expand the inner product by definition

where we recall hT−1=BT∇w^∗fh_{T-1}=B_{T}\nabla_{\hat{w}^{*}}f as ∇λf=0\nabla_{\lambda}f=0 by assumption.

Next we show a technical lemma, which provides a critical tool to bound the second term above; its proof is given in the next section.

Let gg be α\alpha-strongly convex and β\beta-smooth. Assume BtB_{t} and AtA_{t} are Lipschitz continuous in ww, and assume BTB_{T} has full column rank. For γ≤1β\gamma\leq\frac{1}{\beta},

and BT⊤BTB_{T}^{\top}B_{T} is non-singular by assumption,

for some c>0c>0, when TT is large enough and γ\gamma is small enough. The implication holds because ∥dλf∥≤O(∥∇w^∗f∥)\|\text{d}_{\lambda}f\|\leq O(\|\nabla_{\hat{w}^{*}}f\|). ∎

Let CAC_{A} and CBC_{B} be the Lipschitz constant of AtA_{t} and BtB_{t}. First, we see that the inner product can be lower bounded by the following terms

The above lower bounds can be shown by the following inequalities:

Next we upper bound the error terms: Δ1\Delta_{1}, Δ2\Delta_{2}, and Δ3\Delta_{3}. We will use the fact that gradient descent converges linearly when optimizing a strongly convex and smooth function .

Let w0w_{0} be the initial condition. Running gradient descent to optimize an α\alpha-strongly convex and β\beta-smooth function gg, with step size 0<γ≤1β0<\gamma\leq\frac{1}{\beta}, generates a sequence {wt}\{w_{t}\} satisfying

where D=∥w0−w∗∥D=\|w_{0}-w^{*}\| and w∗=arg min⁡g(w)w^{*}=\operatorname*{arg\,min}g(w).

Lemma B.2 implies for T≥tT\geq t, ∥wT−wt∥≤2De−αγt\|w_{T}-w_{t}\|\leq 2De^{-\alpha\gamma t}.

Now we proceed to bound the errors Δ1\Delta_{1}, Δ2\Delta_{2}, and Δ3\Delta_{3}.

Using the bounds on Δ1\Delta_{1}, Δ2\Delta_{2}, and Δ3\Delta_{3}, we prove the final result.

Appendix C Proof of Theorem 3.3

The proof of this theorem is a standard proof of non-convex optimization with biased gradient estimates. Here we include it for completeness, as part of it will be used later in the proof of Theorem 3.4.

Let λτ\lambda_{\tau} be the τ\tauth iterate. For short hand, we write dλf(τ)=dλf(λτ)\text{d}_{\lambda}f_{(\tau)}=\text{d}_{\lambda}f(\lambda_{\tau}), and hT−K,(τ)=hT−K(λτ)h_{T-K,(\tau)}=h_{T-K}(\lambda_{\tau}). Assume FF is LL-smooth and ∥dλf(τ)∥≤G\|\text{d}_{\lambda}f_{(\tau)}\|\leq G and ∥hT−K,(τ)∥≤G\|h_{T-K,(\tau)}\|\leq G almost surely for all τ\tau. Then by LL-smoothness, it satisfies

Let eτ=dλf(τ)−hT−K,(τ)e_{\tau}=\text{d}_{\lambda}f_{(\tau)}-h_{T-K,(\tau)} be the error in the gradient estimate. Substitute the recursive update λτ+1=λτ−ηthT−K,(τ)\lambda_{\tau+1}=\lambda_{\tau}-\eta_{t}h_{T-K,(\tau)} to the above inequality. Conditioned on λτ\lambda_{\tau}, it satisfies

Performing telescoping sum with the above inequality, we have

Dividing both sides by ∑τ=1Rητ\sum_{\tau=1}^{R}\eta_{\tau} and using the facts that ητ=O(1τ)\eta_{\tau}=O(\frac{1}{\sqrt{\tau}}) and that

Appendix D Proof of Theorem 3.4

First we consider the special case when SS is deterministic. Let H≥KH\geq K. We decompose the full gradients into four parts

We assume that wtw_{t} enters a locally strongly convex region for t≥Ht\geq H. This implies, by Proposition 3.1, that ∥e∥≤O(e−αγH∥∇w^∗f∥)\|e\|\leq O(e^{-\alpha\gamma H}\|\nabla_{\hat{w}^{*}}f\|).

To prove the theorem, we first verify two conditions:

By Lemma 3.2, the assumption ∇λf⊤(dλf+hT−K−∇λf)≥Ω(∥∇λf∥2)\nabla_{\lambda}f^{\top}(\text{d}_{\lambda}f+h_{T-K}-\nabla_{\lambda}f)\geq\Omega(\|\nabla_{\lambda}f\|^{2}), and ∥e∥≤O(e−αγH∥∇w^∗f∥)\|e\|\leq O(e^{-\alpha\gamma H}\|\nabla_{\hat{w}^{*}}f\|):

Therefore, for HH large enough, it holds that

By definition of hT−K=∇λf+qh_{T-K}=\nabla_{\lambda}f+q, it holds that

Let ff be a lower-bound and LL-smooth function. Consider the iterative update rule

where gtg_{t} satisfies gt⊤∇f(xt)≥c1ht2g_{t}^{\top}\nabla f(x_{t})\geq c_{1}h_{t}^{2} and ∥gt∥2≤c2ht2\|g_{t}\|^{2}\leq c_{2}h_{t}^{2}, for some constant c1,c2>0c_{1},c_{2}>0 and scalar hth_{t}. Suppose ff is lower-bounded and η\eta is chosen such that (−c1η+Lc2η22)≤0\left(-c_{1}\eta+\frac{Lc_{2}\eta^{2}}{2}\right)\leq 0. Then lim⁡t→∞ht=0\lim\limits_{t\to\infty}h_{t}=0.

By telescoping sum, we can show ∑t=0∞(cη−Lη22)ht2<∞\sum_{t=0}^{\infty}\left(c\eta-\frac{L\eta^{2}}{2}\right)h_{t}^{2}<\infty, which implies lim⁡t→∞ht=0\lim_{t\to\infty}h_{t}=0. ∎

Finally, we prove the main theorem by applying Lemma D.1. Consider a deterministic problem. Take ht2=∥∇λf(λt)∥2+∥∇w^∗f(λt)∥2h_{t}^{2}=\|\nabla_{\lambda}f(\lambda_{t})\|^{2}+\|\nabla_{\hat{w}^{*}}f(\lambda_{t})\|^{2}. Because of (15) and (16), by Lemma D.1, it satisfies that

As ∥dλf∥≤O(∥∇λf∥+∥∇w^∗f∥)\|\text{d}_{\lambda}f\|\leq O(\|\nabla_{\lambda}f\|+\|\nabla_{\hat{w}^{*}}f\|), it shows ∥dλf∥\|\text{d}_{\lambda}f\| converges to zero in the limit.

Appendix E Proof of Theorem 3.5

We prove the non-convergence using the following strategy. First we show that, when assumption 3 in Theorem 3.4, i.e.

does not hold, there is some problem such that hT−k≠0h_{T-k}\neq 0 for all stationary points (i.e. λ\lambda such that dλf=0\text{d}_{\lambda}f=0). Then we show that, for such a problem, optimizing λ\lambda with hT−kh_{T-k} cannot converge to any of the stationary points.

To construct the counterexample, we consider a scalar deterministic bilevel optimization problem of the form

in which ϕ\phi is some perturbation function that we will later define, and w^∗\hat{w}^{*} is computed by performing T>1T>1 steps of gradient descent in the lower-level optimization problem with some constant initial condition w0w_{0} and constant step size 0<γ<10<\gamma<1, i.e.

We can observe this problem satisfies almost all the assumptions in Theorem 3.4:

The lower-level objective gg is smooth and strongly convex. (Proposition 3.1)

The upper-level objective FF is smooth. (Theorem 3.3)

The lower-level objective gg is second-order continuously differentiable (assumption 1 in Theorem 3.4)

The Jacobian if full rank, i.e. Bt=γ>0B_{t}=\gamma>0 (assumption 2 in Theorem 3.4)

The upper-level objective function is deterministic, i.e. F=fF=f (assumption 4 in Theorem 3.4)

But we will show that properly setting ϕ\phi can break the non-interfering assumption in (17) (i.e. assumption 3 in Theorem 3.4) and then creates a problem such that optimizing λ\lambda with KK-RMD does not converge to an exact stationary point.

We follow the two-step strategy mentioned above.

Without loss of generality, let us consider optimizing λ\lambda with 11-RMD. In this case we can write the approximate and the exact gradients in closed form as

which are given by (5) and (8). We will show that by properly choosing ϕ\phi, we can define f(λ)=12(w^∗)2+ϕ(λ)f(\lambda)=\frac{1}{2}(\hat{w}^{*})^{2}+\phi(\lambda) such that, at any of the stationary points of ff, the approximate gradient of 11-RMD does not vanish. That is, we show when dλf=0\text{d}_{\lambda}f=0, hT−1≠0h_{T-1}\neq 0.

Before proceeding, let us define u=w∗γu=w^{*}\gamma and v=w∗γ∑t=0T(1−γ)T−tv=w^{*}\gamma\sum_{t=0}^{T}(1-\gamma)^{T-t} for convenience. To show how to construct ϕ\phi, let us consider the stationary points in the caseNote in this special case, assumption 3 in Theorem 3.4 holds trivially when ϕ(λ)=0\phi(\lambda)=0 (i.e. ∇λf=0\nabla_{\lambda}f=0) and optimizing λ\lambda with KK-RMD converges to an exact stationary point. when ϕ=0\phi=0. Let P0P_{0} denote the set of these stationary points, i.e. P0={λ:v=0}P_{0}=\{\lambda:v=0\}. Since ff is smooth and lower-bounded, we know that P0P_{0} is non-empty, and from the construction of our counterexample we know that P0P_{0} contains exactly the λ\lambdas such that w∗=0w^{*}=0.

We use this fact to pick an adversarial ϕ\phi. Consider any smooth, lower-bounded ϕ\phi whose stationary points are not in P0P_{0}, e.g. ϕ(λ)=12(λ−λ0)2\phi(\lambda)=\frac{1}{2}(\lambda-\lambda_{0})^{2} and λ0∉P0\lambda_{0}\notin P_{0}. Then f(λ)=12(w^∗)2+ϕ(λ)f(\lambda)=\frac{1}{2}(\hat{w}^{*})^{2}+\phi(\lambda) has a non-empty set of stationary points PϕP_{\phi} such that Pϕ∩P0=∅P_{\phi}\cap P_{0}=\emptyset. We see that, for such ϕ\phi, the non-interfering assumption (assumption 3 in Theorem 3.4) is violated in PϕP_{\phi}:

And we show for any λ∈Pϕ\lambda\in P_{\phi} it holds that hT−1≠0h_{T-1}\neq 0. This can be seen from the definition

where the last inequality is because w∗≠0w^{*}\neq 0 for λ∈Pϕ\lambda\in P_{\phi}.

We have shown that there is a problem which satisfies all the assumptions but assumption 3 of Theorem 3.4, and at any of its stationary points (i.e. when dλf=0\text{d}_{\lambda}f=0) we have hT−K≠0h_{T-K}\neq 0. Now we show this property implies failure to converge to the stationary points for the general problems considered in Theorem 3.5 (i.e. we do not rely on the form made in Step 1 anymore).

We prove this by contradiction. Let λ∗\lambda^{*} be one of the stationary points. We choose δ0>0\delta_{0}>0 such that, for some ϵ>0\epsilon>0, ∥hT−K∥>ϵ/γ\|h_{T-K}\|>\epsilon/\gamma for all λ\lambda inside the neighborhood {λ:∥λ−λ∗∥<δ02}\{\lambda:\|\lambda-\lambda^{*}\|<\frac{\delta_{0}}{2}\}, where we recall γ\gamma is the step size of the lower-level optimization problem. A non-zero δ0\delta_{0} exists because hT−1h_{T-1} is continuous by our assumption and hT−K≠0h_{T-K}\neq 0 at λ∗\lambda^{*}.

We are ready to show the contradiction. Let δ=min⁡{δ0,ϵ}\delta=\min\{\delta_{0},\epsilon\}. Suppose there is a sequence {λτ}\{\lambda_{\tau}\} that converges to the stationary point λ∗\lambda^{*}. This means that there is 0<M<∞0<M<\infty such that, ∀τ≥M\forall\tau\geq M, ∥λτ−λ∗∥<δ2\|\lambda_{\tau}-\lambda^{*}\|<\frac{\delta}{2}, which implies that ∀τ≥M\forall\tau\geq M, ∥λτ+1−λτ∥<δ\|\lambda_{\tau+1}-\lambda_{\tau}\|<\delta. However, by our choice of δ0\delta_{0}, ∥λτ+1−λτ∥=γ∥hT−K∥>ϵ≥δ\|\lambda_{\tau+1}-\lambda_{\tau}\|=\gamma\|h_{T-K}\|>\epsilon\geq\delta, leading to a contradiction.

Thus, no sequence {λτ}\{\lambda_{\tau}\} converges to any of the stationary points. This concludes our proof. ∎

Appendix F Proof of Proposition 3.6

Recall our shorthand that ∇λ,wg\nabla_{\lambda,w}g and ∇w,wg\nabla_{w,w}g are evaluated at (w∗,λ)(w^{*},\lambda). In the limit, it holds that

To prove the equality (12), we use Lemma (F.1).

For a matrix AA with ∥A∥<1\|A\|<1, it satisfies that

Since γ≤1β\gamma\leq\frac{1}{\beta}, we have γαI⪯γ∇w,wg⪯I\gamma\alpha I\preceq\gamma\nabla_{w,w}g\preceq I, so ∥I−γ∇w,wg∥<1\|I-\gamma\nabla_{w,w}g\|<1. By Lemma F.1,

Appendix G Detailed experimental setup

In this appendix, we provide more details about the settings we used in each experiment. We use Adam to optimize the upper-level objective and vanilla gradient descent for the lower objective. We denote by w^∗\hat{w}^{*} the results of running TT steps of gradient descent with step size γ\gamma.

In this appendix, we provide more details about the data hypercleaning experiment on MNIST from Section 4.2.1.

Both the training and the validation sets consist of 5000 class-balanced examples from the MNIST dataset. The test set consists of the remaining examples. For each training example, with probability 12\frac{1}{2}, we replaced the label with a uniformly random one.

For various KK, we performed KK-RMD for 1000 hyperiterations. Like in the toy experiment (Section 4.1) we adjusted the initial meta-learning rate η0\eta_{0} for each KK so that the norm of the initial update was roughly the same for each KK.

We asserted earlier that the reported F1 scores are not sensitive to our choice of threshold λi<−3\lambda_{i}<-3. To validate this assertion, we repeated the experiment for various thresholds. F1 scores are reported in the table below.

We only ran these experiments for 150150 hyperiterations, because the F1 score has essentially converged by that point. Indeed, the plot below shows identification of corrupted labels for K=1K=1, with cutoff λi<−4\lambda_{i}<-4. The X axis is in units of 1000 hyperiterations. We see that 11-RMD rapidly identifies most of the mislabeled examples, with a few false positives.

G.2 Task interaction

We use T=100T=100 iterations of gradient descent with learning rate 0.10.1 in the lower objective which yields w^S∗\hat{w}_{S}^{*}. To ensure that CC is symmetric, and that CijC_{ij} and ρ\rho are nonnegative, we re-parametrize them as ρ=softplus(ν)\rho=\text{softplus}(\nu) and C=A+A⊤C=A+A^{\top}, where Aij=softplus(Bij)A_{ij}=\text{softplus}(B_{ij}) and BB is a hyperparameter matrix. Thus, the hyperparameters to be optimized are λ={B,ν}\lambda=\{B,\nu\}.

Rather than using raw pixels, we extract image features from the output of the average pooling layer in Resnet-1818 which is trained on ImageNet . We use the same data pre-processing that is used for training Resnet architecture.

When reporting test accuracy, we run 10 independent trials. In each trial, we sample the training and validation datasets with a balanced set of mm examples each (m=50m=50 for CIFAR-10 and m=300m=300 for CIFAR-100) and use the rest of the dataset for testing. To avoid over-fitting, we use early stopping when the testing error does not improve for 500500 hyper-iterations.

Although we are using a similar setting as Franceschi et al. , our results on full back-propagation are quite different from theirs. We believe it is because we are using a different network architecture and pre-processing method for feature extraction.

G.3 One-shot classification

The Omniglot dataset , a popular benchmark for few-shot learning, is used in this experiment. We consider 55-way classification with 11 training and 1515 validation examples for each of the five classes. To evaluate the generalization performance, we restrict the meta-training dataset to a random subset of 12001200 of the 16231623 Omniglot characters. The meta-validation dataset consists of 100100 other characters, and meta-testing dataset has the remaining 323323 characters. We use the meta-validation dataset for tuning the upper-level optimization parameters and report the performance of the algorithm on the meta-testing dataset. Note that no data augmentation method is used in the training.

The overall neural network architecture is shown in Figure 8. Our architecture inherits the hyper-representation model of Franceschi et al. with some modifications. The first two convolutional layers, parametrized by hyperparameter λ={λl1,λl2}\lambda=\{\lambda_{l_{1}},\lambda_{l_{2}}\}, transform the input image into a “hyper-representation” space. The last three layers, parametrized by w={wl3,wl4,wl5}w=\{w_{l_{3}},w_{l_{4}},w_{l_{5}}\} are fine-tuned in the lower-level optimization. Additionally, we have regularization hyperparameters λr={ρi}i=13∪{cj}j=13\lambda_{r}=\{\rho_{i}\}_{i=1}^{3}\cup\{c_{j}\}_{j=1}^{3}. The overall setup corresponds essentially to meta-learning the two bottom layers of a CNN; for each task, the weights in the first two layers are frozen, and the kk-way classifier of the last three layers is fine tuned. Overall, the model has ≈110\approx 110k hyperparameters and ≈75\approx 75k parameters.

We use a meta-batch-size of 44 in each hyper-iteration. To limit the training time, we stop all the algorithms after 50005000 hyper-iterations. Needless to say, these results could be further improved by using data augmentation, higher meta-batch size, and running more hyper-iterations. However, our current setup is selected so that all the experiments can be run in a reasonable amount of time, while sharing a similar setting used in practical one-shot learning.