Optimizing Millions of Hyperparameters by Implicit Differentiation
Jonathan Lorraine, Paul Vicol, David Duvenaud
Introduction
The generalization of neural networks (NN s) depends crucially on the choice of hyperparameters. Hyperparameter optimization (HO) has a rich history , and achieved recent success in scaling due to gradient-based optimizers . There are dozens of regularization techniques to combine in deep learning, and each may have multiple hyperparameters . If we can scale HO to have as many—or more—hyperparameters as parameters, there are various exciting regularization strategies to investigate. For example, we could learn a distilled dataset with a hyperparameter for every feature of each input , weights on each loss term , or augmentation on each input .
When the hyperparameters are low-dimensional—e.g., - dimensions—simple methods, like random search, work; however, these break down for medium-dimensional HO—e.g., - dimensions. We may use more scalable algorithms like Bayesian Optimization , but this often breaks down for high-dimensional HO—e.g., > dimensions. We can solve high-dimensional HO problems locally with gradient-based optimizers, but this is difficult because we must differentiate through the optimized weights as a function of the hyperparameters. In other words, we must approximate the Jacobian of the best-response function of the parameters to the hyperparameters.
We leverage the Implicit Function Theorem (IFT) to compute the optimized validation loss gradient with respect to the hyperparameters—hereafter denoted the hypergradient. The IFT requires inverting the training Hessian with respect to the NN weights, which is infeasible for modern, deep networks. Thus, we propose an approximate inverse, motivated by a link to unrolled differentiation that scales to Hessians of large NNs, is more stable than conjugate gradient , and only requires a constant amount of memory.
Finally, when fitting many parameters, the amount of data can limit generalization. There are ad hoc rules for partitioning data into training and validation sets—e.g., using % for validation. Often, practitioners re-train their models from scratch on the combined training and validation partitions with optimized hyperparameters, which can provide marginal test-time performance increases. We verify empirically that standard partitioning and re-training procedures perform well when fitting few hyperparameters, but break down when fitting many. When fitting many hyperparameters, we need a large validation partition, which makes re-training our model with optimized hyperparameters vital for strong test performance.
We propose a stable inverse Hessian approximation with constant memory cost.
We show that the IFT is the limit of differentiating through optimization.
We scale IFT-based hyperparameter optimization to modern, large neural architectures, including AlexNet and LSTM-based language models.
We demonstrate several uses for fitting hyperparameters almost as easily as weights, including per-parameter regularization, data distillation, and learned-from-scratch data augmentation methods.
We explore how training-validation splits should change when tuning many hyperparameters.
Overview of Proposed Algorithm
There are four essential components to understanding our proposed algorithm. Further background is provided in Appendix A, and notation is shown in Table 5.
1. HO is nested optimization: Let and denote the training and validation losses, w the NN weights, and the hyperparameters. We aim to find optimal hyperparameters such that the NN minimizes the validation loss after training:
Our implicit function is , which is the best-response of the weights to the hyperparameters. We assume unique solutions to for simplicity.
2. Hypergradients have two terms: For gradient-based HO we want the hypergradient , which decomposes into:
The direct gradient is easy to compute, but the indirect gradient is difficult to compute because we must account for how the optimal weights change with respect to the hyperparameters (i.e., ). In HO the direct gradient is often identically , necessitating an approximation of the indirect gradient to make any progress (visualized in Fig. 1).
3. We can estimate the implicit best-response with the IFT: We approximate the best-response Jacobian—how the optimal weights change with respect to the hyperparameters—using the IFT (Thm. IFT). We present the complete statement in Appendix Theorem, but highlight the key assumptions and results here.
If for some and regularity conditions are satisfied, then surrounding there is a function s.t. and we have:
The condition is equivalent to being a fixed point of the training gradient field. Since is a fixed point of the training gradient field, we can leverage the IFT to evaluate the best-response Jacobian locally. We only have access to an approximation of the true best-response—denoted —which we can find with gradient descent.
4. Tractable inverse Hessian approximations: To exactly invert a general Hessian, we often require operations, which is intractable for the matrix in Eq. IFT in modern NNs . We can efficiently approximate the inverse with the Neumann series:
In Section 4 we show that unrolling differentiation for steps around locally optimal weights is equivalent to approximating the inverse with the first terms in the Neumann series. We then show how to use this approximation without instantiating any matrices by using efficient vector-Jacobian products.
We outline our method in Algs. 1, 2, and 3, where denotes the learning rate. Alg. 3 is also shown in . We visualize the hypergradient computation in Fig. 2.
Related Work
Implicit Function Theorem. The IFT has been used for optimization in nested optimization problems , backpropagating through arbitrarily long RNNs , or even efficient -fold cross-validation . Early work applied the IFT to regularization by explicitly computing the Hessian (or Gauss-Newton) inverse . In , the identity matrix is used to approximate the inverse Hessian in the IFT. HOAG uses conjugate gradient (CG) to invert the Hessian approximately and provides convergence results given tolerances on the optimal parameter and inverse. In iMAML , a center to the weights is fit to perform well on multiple tasks—contrasted with our use of validation loss. In DEQ , implicit differentiation is used to add differentiable fixed-point methods into NN architectures. We use a Neumann approximation for the inverse-Hessian, instead of CG or the identity.
Approximate inversion algorithms. CG is difficult to scale to modern, deep NNs . We use the Neumann inverse approximation, which was observed to be a stable alternative to CG in NNs . The stability is motivated by connections between the Neumann approximation and unrolled differentiation . Alternatively, we could use prior knowledge about the NN structure to aid in the inversion—e.g., by using KFAC . It is possible to approximate the Hessian with the Gauss-Newton matrix or Fisher Information matrix . Various works use an identity approximation to the inverse, which is equivalent to -step unrolled differentiation .
Unrolled differentiation for HO . A key difficulty in nested optimization is approximating how the optimized inner parameters (i.e., NN weights) change with respect to the outer parameters (i.e., hyperparameters). We often optimize the inner parameters with gradient descent, so we can simply differentiate through this optimization. Differentiation through optimization has been applied to nested optimization problems by , was scaled to HO for NNs by , and has been applied to various applications like learning optimizers . provides convergence results for this class of algorithms, while discusses forward- and reverse-mode variants.
As the number of gradient steps we backpropagate through increases, so does the memory and computational cost. Often, gradient descent does not exactly minimize our objective after a finite number of steps—it only approaches a local minimum. Thus, to see how the hyperparameters affect the local minima, we may have to unroll the optimization infeasibly far. Unrolling a small number of steps can be crucial for performance but may induce bias . discusses connections between unrolling and the IFT, and proposes to unroll only the last -steps. DrMAD proposes an interpolation scheme to save memory.
We compare hypergradient approximations in Table 1, and memory costs of gradient-based HO methods in Table 2. We survey gradient-free HO in Appendix B.
Method
In this section, we discuss how HO is a uniquely challenging nested optimization problem and how to combine the benefits of the IFT and unrolled differentiation.
Eq. 3 shows that the hypergradient decomposes into a direct and indirect gradient. The bottleneck in hypergradient computation is usually finding the indirect gradient because we must take into account how the optimized parameters vary with respect to the hyperparameters. A simple optimization approach is to neglect the indirect gradient and only use the direct gradient. This can be useful in zero-sum games like GANs because they always have a non-zero direct term.
However, using only the direct gradient does not work in general games . In particular, it does not work for HO because the direct gradient is identically when the hyperparameters can only influence the validation loss by changing the optimized weights . For example, if we use regularization like weight decay when computing the training loss, but not the validation loss, then the direct gradient is always .
If the direct gradient is identically , we call the game pure-response. Pure-response games are uniquely difficult nested optimization problems for gradient-based optimization because we cannot use simple algorithms that rely on the direct gradient like simultaneous SGD. Thus, we must approximate the indirect gradient.
2 Unrolled Optimization and the IFT
Here, we discuss the relationship between the IFT and differentiation through optimization. Specifically, we (1) introduce the recurrence relation that arises when we unroll SGD optimization, (2) give a formula for the derivative of the recurrence, and (3) establish conditions for the recurrence to converge. Notably, we show that the fixed points of the recurrence recover the IFT solution. We use these results to motivate a computationally tractable approximation scheme to the IFT solution. We give proofs of all results in Appendix D.
Unrolling SGD optimization—given an initialization —gives us the recurrence:
In our exposition, assume that . We provide a formula for the derivative of the recurrence, to show that it converges to the IFT under some conditions.
Given the recurrence from unrolling SGD optimization in Eq. 5, we have:
This recurrence converges to a fixed point if the transition Jacobian is contractive, by the Banach Fixed-Point Theorem . Theorem 8 shows that the recurrence converges to the IFT if we start at locally optimal weights , and the transition Jacobian is contractive. We leverage that if an operator is contractive, then the Neumann series .
Given the recurrence from unrolling SGD optimization in Eq. 5, if :
and if is contractive:
This result is also shown in , but they use a different approximation for computing the hypergradient—see Table 1. Instead, we use the following best-response Jacobian approximation, where controls the trade-off between computation and error bounds:
Shaban et al. use an approximation that scales memory linearly in , while ours is constant. We save memory because we reuse last w times, while needs the last w’s. Scaling the Hessian by the learning rate is key for convergence. Our algorithm has the following main advantages relative to other approaches:
It requires a constant amount of memory, unlike other unrolled differentiation methods .
It is more stable than conjugate gradient, like unrolled differentiation methods .
3 Scope and Limitations
We need continuous hyperparameters to use gradient-based optimization, but many discrete hyperparameters (e.g., number of hidden units) have continuous relaxations . Also, we can only optimize hyperparameters that change the loss manifold, so our approach is not straightforwardly applicable to optimizer hyperparameters.
To exactly compute hypergradients, we must find s.t., which we can only solve to a tolerance with an approximate solution denoted . shows results for error in and the inversion.
Experiments
We first compare the properties of Neumann inverse approximations and conjugate gradient, with experiments similar to . Then we demonstrate that our proposed approach can overfit the validation data with small training and validation sets. Finally, we apply our approach to high-dimensional HO tasks: (1) dataset distillation; (2) learning a data augmentation network; and (3) tuning regularization parameters for an LSTM language model.
HO algorithms that are not based on implicit differentiation or differentiation through optimization—such as —do not scale to the high-dimensional hyperparameters we use. Thus, we cannot sensibly compare to them for high-dimensional problems.
In Fig. 3 we investigate how close various approximations are to the true inverse. We calculate the distance between the approximate hypergradient and the true hypergradient. We can only do this for small-scale problems because we need the exact inverse for the true hypergradient. Thus, we use a linear network on the Boston housing dataset , which makes finding the best-response and inverse training Hessian feasible.
In Fig. 4 we show the inverse Hessian for a fully-connected 1-layer NN on the Boston housing dataset. The true inverse Hessian has a dominant diagonal, motivating identity approximations, while using more Neumann terms yields structure closer to the true inverse.
2 Overfitting a Small Validation Set
In Fig. 5, we check the capacity of our HO algorithm to overfit the validation dataset. We use the same restricted dataset as in of training and validation examples, which allows us to assess HO performance easily. We tune a separate weight decay hyperparameter for each NN parameter as in . We show the performance with a linear classifier, AlexNet , and ResNet . For AlexNet, this yields more than hyperparameters, so we can perfectly classify our validation data by optimizing the hyperparameters.
3 Dataset Distillation
Dataset distillation aims to learn a small, synthetic training dataset from scratch, that condenses the knowledge contained in the original full-sized training set. The goal is that a model trained on the synthetic data generalizes to the original validation and test sets. Distillation is an interesting benchmark for HO as it allows us to introduce tens of thousands of hyperparameters, and visually inspect what is learned: here, every pixel value in each synthetic training example is a hyperparameter. We distill MNIST and CIFAR-/ , yielding $\!\times\!\!\times\!=\!\times\!\!\times\!\!\times\!=\!\times\!\!\times\!\!\times\!=10100$.
4 Learned Data Augmentation
Data augmentation is a simple way to introduce invariances to a model—such as scale or contrast invariance—that improve generalization . Taking advantage of the ability to optimize many hyperparameters, we learn data augmentation from scratch (Fig. 7).
Results for the identity and Neumann inverse approximations are shown in Table 3. We omit CG because it performed no better than the identity. We found that using the data augmentation network improves validation and test accuracy by 2-3%, and yields smaller variance between multiple random restarts. In , a different augmentation network architecture is learned with adversarial training.
5 RNN Hyperparameter Optimization
We also used our proposed algorithm to tune regularization hyperparameters for an LSTM trained on the Penn TreeBank (PTB) corpus . As in , we used a -layer LSTM with hidden units per layer and -dimensional word embeddings. Additional details are provided in Appendix E.4.
Overfitting Validation Data. We first verify that our algorithm can overfit the validation set in a small-data setting with 10 training and 10 validation sequences (Fig. 8). The LSTM architecture we use has weights, and we tune a separate weight decay hyperparameter per weight. We overfit the validation set, reaching nearly 0 validation loss.
Large-Scale HO. There are various forms of regularization used for training RNNs, including variational dropout on the input, hidden state, and output; embedding dropout that sets rows of the embedding matrix to 0, removing tokens from all sequences in a mini-batch; DropConnect on the hidden-to-hidden weights; and activation and temporal activation regularization. We tune these hyperparameters simultaneously. Additionally, we experiment with tuning separate dropout/DropConnect rate for each activation/weight, giving total hyperparameters. To allow for gradient-based optimization of dropout rates, we use concrete dropout .
Instead of using the small dropout initialization as in , we use a larger initialization of , which prevents early learning rate decay for our method. The results for our new initialization with no HO, our method tuning the same hyperparameters as (“Ours”), and our method tuning many more hyperparameters (“Ours, Many”) are shown in Table 4. We are able to tune hyperparameters more quickly and achieve better perplexities than the alternatives.
6 Effects of Many Hyperparameters
Given the ability to tune high-dimensional hyperparameters and the potential risk of overfitting to the validation set, should we reconsider how our training and validation splits are structured? Do the same heuristics apply as for low-dimensional hyperparameters (e.g., use of the data for validation)?
In Fig. 9 we see how splitting our data into training and validation sets of different ratios affects test performance. We show the results of jointly optimizing the NN weights and hyperparameters, as well as the results of fixing the final optimized hyperparameters and re-training the NN weights from scratch, which is a common technique for boosting performance .
We evaluate a high-dimensional regime with a separate weight decay hyperparameter per NN parameter, and a low-dimensional regime with a single, global weight decay. We observe that: (1) for few hyperparameters, the optimal combination of validation data and hyperparameters has similar test performance with and without re-training, because the optimal amount of validation data is small; and (2) for many hyperparameters, the optimal combination of validation data and hyperparameters is significantly affected by re-training, because the optimal amount of validation data needs to be large to fit our hyperparameters effectively.
For few hyperparameters, our results agree with the standard practice of using % of the data for validation and the other % for training. For many hyperparameters, our results show that we should use larger validation partitions for HO . If we use a large validation partition to fit the hyperparameters, it is critical to re-train our model with all of the data.
Conclusion
We present a gradient-based hyperparameter optimization algorithm that scales to high-dimensional hyperparameters for modern, deep NNs . We use the implicit function theorem to formulate the hypergradient as a matrix equation, whose bottleneck is inverting the Hessian of the training loss with respect to the NN parameters. We scale the hypergradient computation to large NNs by approximately inverting the Hessian, leveraging a relationship with unrolled differentiation.
We believe algorithms of this nature provide a path for practical nested optimization, where we have Hessians with known structure. Examples of this include GANs , and other multi-agent games .
Acknowledgements
We thank Chris Pal for recommending we investigate re-training with all the data, Haoping Xu for discussing related experiments on inverse-approximation variants, Roger Grosse for guidance, and Cem Anil & Chris Cremer for their feedback on the paper. Paul Vicol was supported by a JP Morgan AI Fellowship. We also thank everyone else at Vector for helpful discussions and feedback.
References
Appendix A Extended Background
Due to a limited size dataset , there may be a significant difference between the minimizer of the empirical risk and the population risk. We can estimate this difference by partitioning our dataset into training and validation datasets— . We find the minimizer over the training dataset , and estimate its performance on the population risk by evaluating the empirical risk over the validation dataset . We introduce modifications to the empirical training risk to decrease our population risk, parameterized by . These parameters for generalization are called the hyperparameters. We call the modified empirical training risk our training loss for simplicity and denote it . Our validation empirical risk is called validation loss for simplicity and denoted by . Often the validation loss does not directly depend on the hyperparameters, and we just have .
The population risk is estimated by plugging the training loss minimizer into the validation loss for the estimated population risk . We want our hyperparameters to minimize the estimated population risk: . We can create a third partition of our dataset to assess if we have overfit the validation dataset with our hyperparameters .
Appendix B Extended Related Work
Independent HO : A simple class of HO algorithms involve making a number of independent hyperparameter selections, and training the model to completion on them. Popular examples include grid search and random search . Since each hyperparameter selection is independent, these algorithms are trivial to parallelize.
Global HO : Some HO algorithms attempt to find a globally optimal hyperparameter setting, which can be important if the loss is non-convex. A simple example is random search, while a more sophisticated example is Bayesian optimization . These HO algorithms often involve re-initializing the hyperparameter and weights on each optimization iteration. This allows global optimization, at the cost of expensive re-training weights or hyperparameters.
Local HO : Other HO algorithms only attempt to find a locally optimal hyperparameter setting. Often these algorithms will maintain a current estimate of the best combination of hyperparameter and weights. On each optimization iteration, the hyperparameter is adjusted by a small amount, which allows us to avoid excessive re-training of the weights on each update. This is because the new optimal weights are near the old optimal weights due to a small change in the hyperparameters.
Learned proxy function based HO : Many HO algorithms attempt to learn a proxy function for optimization. The proxy function is used to estimate the loss for a hyperparameter selection. We could learn a proxy function for global or local HO . We can learn a useful proxy function over any node in our computational graph including the optimized weights. For example, we could learn how the optimized weights change w.r.t. the hyperparameters , how the optimized predictions change w.r.t. the hyperparameters , or how the optimized validation loss changes w.r.t. the hyperparameters as in Bayesian Optimization. It is possible to do gradient descent on the proxy function to find new hyperparameters to query as in Bayesian optimization. Alternatively, we could use a non-differentiable proxy function to get cheap estimates of the validation loss like SMASH for architecture choices.
Appendix C Implicit Function Theorem
Let be a continuously differentiable function. Fix a point with . If the Jacobian is invertible, there exists an open set containing s.t. there exists a continuously differentiable function s.t.:
Moreover, the partial derivatives of in are given by the matrix product:
Appendix D Proofs
If the recurrence given by unrolling SGD optimization in Eq. 5 has a fixed point (i.e., ), then:
Given the recurrence from unrolling SGD optimization in Eq. 5 we have:
Given the recurrence from unrolling SGD optimization in Eq. 5, if :
and if is contractive:
Appendix E Experiments
We use PyTorch as our computational framework. All experiments were performed on NVIDIA TITAN Xp GPUs.
For all CNN experiments we use the following optimization setup: for the NN weights we use Adam with a learning rate of 1e-4. For the hyperparameters we use RMSprop with a learning rate of 1e-2.
E.2 Dataset Distillation
With MNIST we use the entire dataset in validation, while for CIFAR we use validation data points.
E.3 Learned Data Augmentation
Augmentation Network Details: Data augmentation can be framed as an image-to-image transformation problem; inspired by this, we use a U-Net as the data augmentation network. To allow for stochastic transformations, we feed in random noise by concatenating a noise channel to the input image, so the resulting input has 4 channels.
E.4 RNN Hyperparameter Optimization
Our base our implementation on the AWD-LSTM codebase https://github.com/salesforce/awd-lstm-lm. Similar to we used a -layer LSTM with hidden units per layer and -dimensional word embeddings.
We used a subset of training sequences and 10 validation sequences, and tuned separate weight decays per parameter. The LSTM architecture we use has weights, and thus an equal number of weight decay hyperparameters.
Optimization Details:
For the large-scale experiments, we follow the training setup proposed in : for the NN weights, we use SGD with learning rate and gradient clipping to magnitude . The learning rate was decayed by a factor of 4 based on the nonmonotonic criterion introduced by (i.e., when the validation loss fails to decrease for 5 epochs). To optimize the hyperparameters, we used Adam with learning rate . We trained on sequences of length 70 in mini-batches of size 40.