Hyperparameter Ensembles for Robustness and Uncertainty Quantification

Florian Wenzel, Jasper Snoek, Dustin Tran, Rodolphe Jenatton

Introduction

Neural networks are well-suited to form ensembles of models . Indeed, neural networks trained from different random initialization can lead to equally well-performing models that are nonetheless diverse in that they make complementary errors on held-out data . This property is explained by the multi-modal nature of their loss landscape and the randomness induced by both their initialization and the stochastic methods commonly used to train them .

Many mechanisms have been proposed to further foster diversity in ensembles of neural networks, e.g., based on cyclical learning rates or Bayesian analysis . In this paper, we focus on exploiting the diversity induced by combining neural networks defined by different hyperparameters. This concept is already well-established and the auto-ML community actively applies it . We build upon this research with the following two complementary goals.

First, for performance independent of computational and memory budget, we seek to improve upon deep ensembles , the current state-of-the-art ensembling method in terms of robustness and uncertainty quantification . To this end, we develop a simple stratification scheme which combines random search and the greedy selection of hyperparameters from with the benefit of multiple random initializations per hyperparameter like in deep ensembles. Figure 1 illustrates our algorithm for a Wide ResNet 28-10 where it leads to substantial improvements, highlighting the benefits of combining different initialization and hyperparameters.

Second, we seek to improve upon batch ensembles , the current state-of-the-art in efficient ensembles. To this end, we propose a parameterization combining that of and self-tuning networks , which enables both weight and hyperparameter diversity. Our approach is a drop-in replacement that outperforms batch ensembles and does not need a separate tuning of the hyperparameters.

Ensembles over neural network weights. Combining the outputs of several neural networks to improve their single performance has a long history, e.g., . Since the quality of an ensemble hinges on the diversity of its members , many mechanisms were developed to generate diverse ensemble members. For instance, cyclical learning-rate schedules can explore several local minima where ensemble members can be snapshot. Other examples are MC dropout or the random initialization itself, possibly combined with the bootstrap . More generally, Bayesian neural networks can be seen as ensembles with members being weighted by the (approximated) posterior distribution over the parameters .

Hyperparameter ensembles. Hyperparameter-tuning methods typically produce a pool of models from which ensembles can be constructed post hoc, e.g., . This idea has been made systematic as part of auto-sklearn and successfully exploited in several other contexts, e.g., and specifically for neural networks as well as in computer vision and genetics . In particular, the greedy ensemble construction from (and later variations thereof ) was shown to work best among other algorithms, either more expensive or more prone to overfitting. To the best of our knowledge, such ensembles based on hyperparameters have not been studied in the light of predictive uncertainty. Moreover, we are not aware of existing methods to efficiently build such ensembles, similarly to what batch ensembles do for deep ensembles. Finally, recent research in Bayesian optimization has also focused on directly optimizing the performance of the ensemble while tuning the hyperparameters .

Hyperparameter ensembles also connect closely to probabilistic models over structures. These works often analyze Bayesian nonparametric distributions, such as over depth and width of a neural network, leveraging Markov chain Monte Carlo for inference . In this work, we examine more parametric assumptions, building on the success of variational inference and mixture distributions: for example, the validation step in hyper-batch ensemble can be viewed as a mixture variational posterior and the entropy penalty is the ELBO’s KL divergence toward a uniform prior.

Concurrent to our paper, construct neural network ensembles within the context of neural architecture search, showing improved robustness for predictions with distributional shift. One of their methods, NES-RS, has similarities with our hyper-deep ensembles (see Section 3), also relying on both random search and to form ensembles, but do not stratify over different initializations. We vary the hyperparameters while keeping the architecture fixed while study the converse. Furthermore, do not explore a parameter- and computationally-efficient method (see Section 4).

Efficient hyperparameter tuning & best-response function. Some hyperparameters of a neural network, e.g., its L2L_{2} regularization parameter(s), can be optimized by estimating the best-response function , i.e., the mapping from the hyperparameters to the parameters of the neural networks solving the problem at hand . Learning this mapping is an instance of learning an hypernetwork and falls within the scope of bilevel optimization problems . Because of the daunting complexity of this mapping, proposed scalable local approximations of the best-response function. Similar methodology was also employed for style transfer and image compression . The self-tuning networks from are an important building block of our approach wherein we extend their setting to the case of an ensemble over different hyperparameters.

2 Contributions

We examine two regimes to exploit hyperparameter diversity: (a) ensemble performance independent of budget and (b) ensemble performance seeking parameter efficiency, where, respectively, deep and batch ensembles are state-of-the-art. We propose one ensemble method for each regime:

(a) Hyper-deep ensembles. We define a greedy algorithm to form ensembles of neural networks exploiting two sources of diversity: varied hyperparameters and random initialization. By stratifying models with respect to the latter, our algorithm subsumes deep ensembles that we outperform in our experiments. Our approach is a simple, strong baseline that we hope will be used in future research.

(b) Hyper-batch ensembles. We efficiently construct ensembles of neural networks defined over different hyperparameters. Both the ensemble members and their hyperparameters are learned end-to-end in a single training procedure, directly maximizing the ensemble performance. Our approach outperforms batch ensembles and generalizes the layer structure of and , while keeping their original memory compactness and efficient minibatching for parallel training and prediction.

We illustrate the benefits of our two ensemble methods on image classification tasks, with multi-layer perceptron, LeNet, ResNet 20 and Wide ResNet 28-10 architectures, in terms of both predictive performance and uncertainty. The code for generic hyper-batch ensemble layers can be found in https://github.com/google/edward2 and the code to reproduce the experiments of Section 5.2 is part of https://github.com/google/uncertainty-baselines.

Background

Deep ensembles are a simple ensembling method where neural networks with different random initialization are combined. Deep ensembles lead to remarkable predictive performance and robust uncertainty estimates . Given some hyperparameters λ0{\boldsymbol{\lambda}}_{0}, a deep ensemble of size KK amounts to solving KK times (1) with random initialization and aggregating the outputs of {fθ^k(λ0)(⋅,λ0)}k=1K\{f_{\hat{{\boldsymbol{\theta}}}_{k}({\boldsymbol{\lambda}}_{0})}(\cdot,{\boldsymbol{\lambda}}_{0})\}_{k=1}^{K}.

2 Self-tuning networks

Since θ′{\boldsymbol{\theta}}^{\prime} captures changes in θ{\boldsymbol{\theta}} induced by changes in λ{\boldsymbol{\lambda}}, replace the typical objective (1), defined for a single value of λ{\boldsymbol{\lambda}}, with an expected objective ,

where p(λ)p({\boldsymbol{\lambda}}) denotes some distribution over the hyperparameters λ{\boldsymbol{\lambda}}. When pp is kept fixed during the optimization of (4), the authors of observed that θ^(λ)\hat{{\boldsymbol{\theta}}}({\boldsymbol{\lambda}}) is not well approximated and proposed instead to use a distribution pt(λ)=p(λ∣ξt)p_{t}({\boldsymbol{\lambda}})=p({\boldsymbol{\lambda}}|{\boldsymbol{\xi}}_{t}) varying with the iteration tt. In our work we choose p(⋅∣ξt)p(\cdot|{\boldsymbol{\xi}}_{t}) to be a log-uniform distribution with ξt{\boldsymbol{\xi}}_{t} containing the bounds of the ranges of λ{\boldsymbol{\lambda}} (see Section 4). The key benefit from (4) is that a single (though, more costly) training gives access to a mapping λ↦fΘ^(⋅,λ){\boldsymbol{\lambda}}\mapsto f_{\hat{{\boldsymbol{\Theta}}}}(\cdot,{\boldsymbol{\lambda}}) which approximates the behavior of fΘ^f_{\hat{{\boldsymbol{\Theta}}}} for hyperparameters in the support of p(λ)p({\boldsymbol{\lambda}}).

The procedure followed by consists in alternating between training and tuning steps. First, the training step performs a stochastic gradient update of Θ{\boldsymbol{\Theta}} in (4), jointly sampling λ∼p(λ∣ξt){\boldsymbol{\lambda}}\sim p({\boldsymbol{\lambda}}|{\boldsymbol{\xi}}_{t}) and (x,y)∈D({\mathbf{x}},y)\in\mathcal{D}. Second, the tuning step makes a stochastic gradient update of ξt{\boldsymbol{\xi}}_{t} by minimizing some validation objective (e.g., the cross entropy):

In (5), derivatives are taken through samples λ∼p(λ∣ξt){\boldsymbol{\lambda}}\sim p({\boldsymbol{\lambda}}|{\boldsymbol{\xi}}_{t}) by applying the reparametrization trick . To prevent p(λ∣ξt)p({\boldsymbol{\lambda}}|{\boldsymbol{\xi}}_{t}) from collapsing to a degenerate distribution, and inspired by variational inference, the authors of add an entropy regularization term H[⋅]\mathcal{H}[\cdot] controlled by τ≥0\tau\geq 0 so that (5) becomes

Hyper-deep ensembles

Figure 2-(left) visualizes different models fθ(⋅,λ)f_{\boldsymbol{\theta}}(\cdot,{\boldsymbol{\lambda}}) according to their hyperparameters λ{\boldsymbol{\lambda}} along the xx-axis and their initialization θinit.{\boldsymbol{\theta}}_{\text{init.}} on the yy-axis. In this view, a deep ensemble corresponds to a “column” where models with different random initialization are combined together, for a fixed λ{\boldsymbol{\lambda}}. On the other hand, a “row” corresponds to the combination of models with different hyperparameters. Such a “row” typically stems from the application of some hyperparameter-tuning techniques .

Given the simplicity, broad applicability, and performance of the greedy algorithm from —e.g., in auto-ML settings , we use it as our canonical procedure to generate a “row”, i.e., an ensemble of neural networks with fixed parameter initialization and various hyperparameters. We refer to it as fixed init hyper ensemble. For completeness, we recall the procedure from in Appendix A (Algorithm 2, named hyper_ens). Given an input set of models (e.g., from random search), hyper_ens greedily grows an ensemble until some target size KK is met by selecting the model with the best improvement of some score, e.g., the validation log-likelihood. We select the models with replacement to be able to learn weighted combinations thereof (see Section 2.1 in ). Note that the procedure from does not require the models to have a fixed initialization: we consider here a fixed initialization to isolate the effect of just varying the hyperparameters (while deep ensembles vary only the initialization, with fixed hyperparameters).

Our goal is two-fold: (a) we want to demonstrate the complementarity of random initialization and hyperparameters as sources of diversity in the ensemble, and (b) design a simple algorithmic scheme that exploits both sources of diversity while encompassing the construction of deep ensembles as a subcase. We defer to Section 5 the study of (a) and next focus on (b).

We proceed in three main steps, as summarized in Algorithm 1. In lines 1-2, we first generate one “row” according to hyper_ens based on the results of random search as input. We then tile and stratify that “row” by training the models for different random initialization (see lines 4-7). The resulting set of models is illustrated in Figure 2-(left). In line 10, we finally re-apply hyper_ens on that stratified set of models to extract an ensemble that can exploit the two sources of diversity. By design, a deep ensemble is one possible outcome of this procedure—one “column”—and so is fixed init hyper ensemble described in the previous paragraph—one “row”.

In lines 1-2, running random search leads to a set of κ\kappa models (i.e., M0\mathcal{M}_{0}). If we were to stratify all of them, we would need KK seeds for each of those κ\kappa models, hence a total of O(κK)\mathcal{O}(\kappa K) models to train. However, we first apply hyper_ens to extract KK models out of the κ\kappa available ones, with K≪κK\ll\kappa. The stratification then needs KK seeds for each of those KK models (lines 4-7), thus O(K2)\mathcal{O}(K^{2}) models to train. We will see in Section 5 that even with standard hyperparameters, e.g., dropout or L2L_{2} parameters, Algorithm 1 can lead to substantial improvements over deep ensembles. In Section C.7.5, we conduct ablation studies to relate to the top-KK strategy used in and NES-RS from .

Hyper-batch ensembles

This section presents our efficient approach to construct ensembles over different hyperparameters.

The core idea lies in the composition of the layers used by batch ensembles for ensembling parameters and self-tuning networks for parameterizing the layer as an explicit function of hyperparameters. The composition preserves complementary features from both approaches.

We continue the example of the dense layer from Section 2.1-Section 2.2. The convolutional layer is described in Section B.1. Assuming an ensemble of size KK, we have for k∈{1,…,K}k\in\{1,\dots,K\}

As noted by , formulation (2) includes a set of rank-1 factors which diversify individual ensemble member weights. In (7), the rank-1 factors rksk⊤{\mathbf{r}}_{k}{\mathbf{s}}_{k}^{\top} and ukvk⊤{\mathbf{u}}_{k}{\mathbf{v}}_{k}^{\top} capture this weight diversity for each respective term.

As noted by , formulation (3) captures local hyperparameter variations in the vicinity of some λ{\boldsymbol{\lambda}}. The term [Δ∘(ukvk⊤)]∘e(λk)⊤[\Delta\circ({\mathbf{u}}_{k}{\mathbf{v}}_{k}^{\top})]\circ{\mathbf{e}}({\boldsymbol{\lambda}}_{k})^{\top} in (7) extends this behavior to the vicinity of the KK hyperparameters {λ1,…,λK}\{{\boldsymbol{\lambda}}_{1},\dots,{\boldsymbol{\lambda}}_{K}\} indexing the KK ensemble members.

Equation (7) maintains the compactness of the original layers of with a resulting memory footprint about twice as large as and equivalent to up to the rank-1 factors.

From an implementation perspective, (7) enables direct reuse of existing code, e.g., DenseBatchEnsemble and Conv2DBatchEnsemble from . The implementation of our layers can be found in https://github.com/google/edward2.

2 Objective function: from single model to ensemble

We first need to slightly overload the notation from Section 2.2 and we write fΘ(x,λk)f_{\boldsymbol{\Theta}}({\mathbf{x}},{\boldsymbol{\lambda}}_{k}) to denote the prediction for the input x{\mathbf{x}} of the kk-th ensemble member indexed by λk{\boldsymbol{\lambda}}_{k}. In Θ{\boldsymbol{\Theta}}, we pack all the parameters of ff, as those described in the example of the dense layer in Section 4.1. In particular, predicting with λk{\boldsymbol{\lambda}}_{k} is understood as using the corresponding parameters {Wk(λk),bk(λk)}\{{\mathbf{W}}_{k}({\boldsymbol{\lambda}}_{k}),{\mathbf{b}}_{k}({\boldsymbol{\lambda}}_{k})\} in (7).

We want the ensemble members to account for a diverse combination of hyperparameters. As a result, each ensemble member is assigned its own distribution of hyperparameters, which we write pt(λk)=p(λk∣ξk,t)p_{t}({\boldsymbol{\lambda}}_{k})=p({\boldsymbol{\lambda}}_{k}|{\boldsymbol{\xi}}_{k,t}) for k∈{1,…,K}k\in\{1,\dots,K\}. Along the line of (4), we consider an expected training objective which now simultaneously operates over ΛK={λk}k=1K{\boldsymbol{\Lambda}}_{K}=\{{\boldsymbol{\lambda}}_{k}\}_{k=1}^{K}

and where L\mathcal{L}, compared with (1), is extended to handle the ensemble predictions

Note that the extensions (8)-(9) with K=1K=1 fall back to the standard formulation of . In our experiments, we take Ω\Omega to be L2L_{2} regularizers applied to the parameters Wk(λk){\mathbf{W}}_{k}({\boldsymbol{\lambda}}_{k}) and bk(λk){\mathbf{b}}_{k}({\boldsymbol{\lambda}}_{k}) of each ensemble member. In Section B.2, we show how to efficiently vectorize the computation of Ω\Omega across the ensemble members and mini-batches of {λk}k=1K\{{\boldsymbol{\lambda}}_{k}\}_{k=1}^{K} sampled from qtq_{t}, as required by (8). In practice, we use one sample of ΛK{\boldsymbol{\Lambda}}_{K} for each data point in the batch: for MLP/LeNet (Section 5.1), we use 256, while for ResNet-20/W. ResNet-28-10 (Section 5.2), we use 512 (64 for each of 8 workers).

Experiments

Throughout the experiments, we use both metrics that depend on the predictive uncertainty—negative log-likelihood (NLL) and expected calibration error (ECE) —and metrics that do not, e.g., the classification accuracy. The supplementary material also reports Brier score (for which we typically observed a strong correlation with NLL). Moreover, as diversity metric, we take the predictive disagreement of the ensemble members normalized by (1-accuracy), as used in . In the tables, we write the number of ensemble members in brackets “(⋅\cdot)” next to the name of the methods.

To validate our approaches and run numerous ablation studies, we first focus on small-scale models, namely MLP and LeNet , over CIFAR-100 and Fashion MNIST . For both models, we add a dropout layer before their last layer. For each pair of dataset/model type, we consider two tuning settings involving the dropout rate and different L2L_{2} regularizers defined with varied granularity, e.g., layerwise. Section C.1 gives all the details about the training, tuning and dataset definitions.

We compare our methods (i) hyper-deep​ ens: hyper-deep ensemble of Section 3 and (ii) hyper-batch​ ens: hyper-batch ensemble of Section 4, to (a) rand​ search: the best single model after 50 trials of random search , (b) Bayes​ opt: the best single model after 50 trials of Bayesian optimization , (c) deep​ ens: deep ensemble using the best hyperparameters found by random search, (d) batch​ ens: batch ensemble , (e) STN: self-tuning networks , and (f) fixed​ init​ hyper​ ens: defined in Section 3. The supplementary material details how we tune the hyperparameters specific to batch​ ens, STN and hyper-batch​ ens (see Section C.2, Section C.3 and Section C.4 and further ablations about e{\mathbf{e}} in Section C.5 and τ\tau in Section C.6). Note that batch​ ens needs the tuning of its own hyperparameters and those of the MLP/LeNet models, while STN and hyper-batch​ ens automatically tune the latter.

We highlight below the key conclusions from Table 1 with single models and ensemble of sizes 3. The same conclusions can also be drawn for the ensemble of size 5 (see Section C.7.1).

With the pictorial view of Figure 2 in mind, fixed​ init​ hyper​ ens, i.e., a “row”, tends to outperform deep​ ens, i.e., a “column”. Moreover, those two approaches (as well as the other methods of the benchmark) are outperformed by our stratified procedure hyper-deep​ ens, demonstrating the benefit of combining hyperparameter and initialization diversity (see Section C.7.2 for the detailed assessment of the statistical significance). In Section C.7.3, we study more specifically the diversity and we show that hyper-deep​ ens has indeed more diverse predictions than deep​ ens.

Among the efficient approaches (the three rightmost columns of Table 1), hyper-batch​ ens performs best. It improves upon both STN and batch​ ens, the two methods it builds upon. In line with , STN typically matches or improves upon rand​ search and Bayes​ opt. As explained in Section 4.1, hyper-batch​ ens has however twice the number of parameters of batch​ ens. In Section C.7.4, we thus compare with a “deep ensemble of two batch ensembles” (i.e., resulting in the same number of parameters but twice as many members as for hyper-batch​ ens). In that case, hyper-batch​ ens also either improves upon or matches the performance of the combination of two batch​ ens.

2 ResNet-20 and Wide ResNet-28-10 on CIFAR-10 & CIFAR-100

We evaluate our approach in a large-scale setting with ResNet-20 and Wide ResNet 28-10 models as they are simple architectures with competitive performance on image classification tasks. We consider six different L2L_{2} regularization hyperparameters (one for each block of the ResNet) and a label smoothing hyperparameter. We show results on CIFAR-10, CIFAR-100 and corruptions on CIFAR-10 . Moreover, in Section D.3, we provide additional out-of-distribution evaluations along the line of . Further details about the experiment settings can be found in Appendix D.

We compare hyper-deep​ ens with a single model (tuned as next explained) and deep​ ens of varying ensemble sizes. Our hyper-deep​ ens is constructed based on 100 trials of random search while deep​ ens and single take the best hyperparameter configuration found by the random search procedure. Figure 1 displays the results on CIFAR-100 along with the standard errors and shows that throughout the ensemble sizes, there is a substantial performance improvement of hyper-deep ensembles over deep ensembles. The results for CIFAR-10 are shown in Appendix D where hyper-deep​ ens leads to consistent but smaller improvements, e.g., in terms of NLL. We next fix the ensemble size to four and compare the performance of hyper-batch​ ens with the direct competing method batch​ ens, as well as with hyper-deep​ ens, deep​ ens and single.

The results are reported in Table 2. On CIFAR-100, hyper-batch​ ens improves, or matches, batch​ ens across all metrics. For instance, in terms of NLL, it improves upon batch​ ens by about 7% and 2% for ResNet-20 and Wide ResNet 28-10 respectively. Moreover, the members of hyper-batch​ ens make more diverse predictions than those of batch​ ens. On CIFAR-10 hyper-batch​ ens also achieves a consistent improvement, though less pronounced (see Table 2). On the same Wide ResNet 28-10 benchmark, with identical training and evaluation pipelines (see https://github.com/google/uncertainty-baselines), variational inference leads to (NLL, ACC, ECE)=(0.211, 0.947, 0.029) and (NLL, ACC, ECE)=(0.944, 0.778, 0.097) for CIFAR-10 and CIFAR-100 respectively, while Monte Carlo dropout gets (NLL, ACC, ECE)=(0.160, 0.959, 0.024) and (NLL, ACC, ECE)=(0.830, 0.776, 0.050) for CIFAR-10 and CIFAR-100 respectively.

We can finally look at how the joint training in hyper-batch​ ens leads to complementary ensemble members. For instance, for Wide ResNet 28-10 on CIFAR-100, while the ensemble performance are (NLL, ACC)=(0.678, 0.820) (see Table 2), the individual members obtain substantially poorer performance, as measured by the average ensemble-member metrics (NLL, ACC)=(0.904, 0.788).

Both in terms of the number of parameters and training time, hyper-batch​ ens is about twice as costly as batch​ ens. For CIFAR-100, hyper-batch​ ens takes 2.16 minutes/epoch and batch​ ens 1.10 minute/epoch. More details are available in Section D.6.

We measure the calibrated prediction on corrupted datasets, which is a type of out-of-distribution examples. We consider the recently published dataset by , which consists of over 30 types of corruptions to the images of CIFAR-10. A similar benchmark can be found in . On Figure 3, we find that all ensembles methods improve upon the single model. The mean accuracies are similar for all ensemble methods, whereas hyper-batch​ ens shows more robustness than batch​ ens as it typically leads to smaller worst values (see bottom whiskers in Figure 3). Plots for calibration error and NLL can be found in Section D.5.

Discussion

We envision several promising directions for future research.

In this work, we have used the layers from that lead to a 2x increase in memory compared with standard layers. In lieu of (3), low-rank parametrizations, e.g., W+∑j=1hej(λ)gjhj⊤{\mathbf{W}}+\sum_{j=1}^{h}e_{j}({\boldsymbol{\lambda}}){\mathbf{g}}_{j}{\mathbf{h}}_{j}^{\top}, would be appealing to reduce the memory footprint of self-tuning networks and hyper-batch ensembles. We formally show in Appendix E that this family of parametrizations is well motivated in the case of shallow models where they enjoy good approximation guarantees.

Our proposed hyperparameter ensembles provide diversity with respect to hyperparameters related to regularization and optimization. We would like to go further in ensembling very different functions in the search space, such as network width, depth , and the choice of residual block. Doing so connects to older work on Bayesian marginalization over structures . More broadly, we can wonder what other types of diversity matter to endow deep learning models with better uncertainty estimates?

Broader Impact

Our work belongs to a broader research effort that tries to quantify the predictive uncertainty for deep neural networks. Those models are known to generalize poorly to small changes to the data while maintaining high confidence in their predictions.

The broader topic of our work is becoming increasingly important in a context where machine learning systems are being deployed in safety-critical fields, e.g., medical diagnosis and self-driving cars . Those examples would benefit from the general technology we contribute to. In those cases, it is essential to be able to reliably trust the uncertainty output by the models before any decision-making process, to possibly escalate uncertain decisions to appropriate human operators.

We are not aware of a group of people that may be put at disadvantage as a result of this direct research.

By definition, our research could contribute to aspects of machine-learning systems used in high-risk domains (e.g., we mentioned earlier medical fields and self-driving cars) which involves complex data-driven decision-making processes. Depending on the nature of the application at hand, a failure of the system could lead to extremely negative consequences. A case in point is the recent screening system used by one third of UK government councils to allocate welfare budget. Link to the corresponding article in The Guardian, October 2019: https://www.theguardian.com/society/2019/oct/15/councils-using-algorithms-make-welfare-decisions-benefits.

The method we develop in this work is domain-agnostic and does not rely on specific data assumptions. Our method also does not contain components that would prevent its combination with existing fairness or privacy-preserving technologies .

Acknowledgments

We would like to thank Nicolas Le Roux, Alexey Dosovitskiy and Josip Djolonga for insightful discussions at earlier stages of this project. Moreover, we would like to thank Sebastian Nowozin, Klaus-Robert Müller and Balaji Lakshminarayanan for helpful comments on a draft of this paper.

References

Supplementary Material: Hyperparameter Ensembles for Robustness and Uncertainty Quantification

Appendix A Further details about fixed init hyper ensembles and hyper-deep ensembles

We recall the procedure from in Algorithm 2. In words, given a pre-defined set of models M\mathcal{M} (e.g., the outcome of random search), we greedily grow an ensemble, until some target size KK is met, by selecting with replacement the model leading to the best improvement of some score S\mathcal{S} such as the validation negative log-likelihood.

The with-replacement selection strategy makes it possible to construct ensembles where the contributions of each member is weighted (see Section 2.1 in ). To properly account for the fact that there may be multiple times the same model selected, we use “.unique()” in Algorithms 1-2 to correctly count the number of members.

Appendix B Further details about hyper-batch ensemble

We detail the structure of the (two-dimensional) convolutional layer of hyper-batch ensemble in the case of KK ensemble members. Similar to the dense layer presented in Section 4.1, the convolutional layer is obtained by composing the layer of batch ensemble and that of self-tuning networks .

where the rank-1 factors are understood to be broadcast along the first two dimensions. Similar, for the bias terms, we have

with δk,e′(λk){\boldsymbol{\delta}}_{k},{\mathbf{e}}^{\prime}({\boldsymbol{\lambda}}_{k}) of the same shape as bk{\mathbf{b}}_{k}.

Given the form of (10) and (11), we can observe that the conclusions drawn for the dense layer in Section 4.1 also hold for the convolutional layer.

We focus on the example of a given dense layer, with weight matrix Wk(λk){\mathbf{W}}_{k}({\boldsymbol{\lambda}}_{k}) and bias term bk(λk){\mathbf{b}}_{k}({\boldsymbol{\lambda}}_{k}), as exposed in Section 4.1.

Let us consider a minibatch of size bb for the KK ensemble members, i.e., {λk,i}k=1K\{{\boldsymbol{\lambda}}_{k,i}\}_{k=1}^{K} for i∈{1,…,b}i\in\{1,\dots,b\}. Moreover, let us introduce the scalar νk,i\nu_{k,i} that is equal to the entry in λk,i{\boldsymbol{\lambda}}_{k,i} containing the value of the L2L_{2} penalty for the particular dense layer under study.The precise relationship between νk,i\nu_{k,i} and λk,i{\boldsymbol{\lambda}}_{k,i} depends on the implementation details and on how the hyperparameters of the problem, e.g., the dropout rates or L2L_{2} penalties, are stored in the vector λk,i{\boldsymbol{\lambda}}_{k,i}.

With that notation, we concentrate on the efficient computation (especially the vectorization with respect to the minibatch dimension) of

the case of the bias term following along the same lines. From Section 4.1 we have

which we have simplified by introducing a few additional shorthands. Let us further introduce

We then develop ∥Wk(λk,i)∥2\|{\mathbf{W}}_{k}({\boldsymbol{\lambda}}_{k,i})\|^{2} into ∥Wk∥2+2Wk⊤(Δk∘ek,i⊤)+∥Δk∘ek,i⊤∥2\|{\mathbf{W}}_{k}\|^{2}+2{\mathbf{W}}_{k}^{\top}({\boldsymbol{\Delta}}_{k}\circ{\mathbf{e}}_{k,i}^{\top})+\|{\boldsymbol{\Delta}}_{k}\circ{\mathbf{e}}_{k,i}^{\top}\|^{2} and plug the decomposition into (12), with Δk2=Δk∘Δk{\boldsymbol{\Delta}}_{k}^{2}={\boldsymbol{\Delta}}_{k}\circ{\boldsymbol{\Delta}}_{k}, leading to

for which all the remaining operations can be efficiently broadcast.

We discuss in this section additional details about the choice of the distributions over the hyperparameters pt(λk)=p(λk∣ξk,t)p_{t}({\boldsymbol{\lambda}}_{k})=p(\lambda_{k}|{\boldsymbol{\xi}}_{k,t}).

In the experiments of Section 5, we manipulate hyperparameters λk{\boldsymbol{\lambda}}_{k}’s that are positive and bounded (e.g., a dropout rate). To simplify the exposition, let us focus momentarily on a single ensemble member (K=1K=1). Let us further consider such a positive, bounded one-dimensional hyperparameter λ∈[a,b]\lambda\in[a,b], with  0<a<b\ 0<a<b, and define ϕ(t)=(b−a) sigmoid(t)+a\phi(t)=(b-a)\ \texttt{sigmoid}(t)+a, with ϕ−1\phi^{-1} its inverse. In that setting, propose to use for pt(λ)=p(λ∣ξt)p_{t}(\lambda)=p(\lambda|{\boldsymbol{\xi}}_{t}) the following distribution:

In preliminary experiments we carried out, we encountered issues with (13), e.g., λ\lambda consistently pushed to its lower bound aa during the optimization.

We have therefore departed from (13) and have focused instead on a simple log-uniform distribution, which is a standard choice for hyperparameter tuning, e.g., . Its probability density function is given by

while its entropy equals H[p(λ∣ξt)]=0.5(log⁡(a)+log⁡(b))+log⁡(log⁡(b/a))\mathcal{H}[p(\lambda|{\boldsymbol{\xi}}_{t})]=0.5(\log(a)+\log(b))+\log(\log(b/a)). The mean of the distribution is given by (b−a)/(log⁡(b)−log⁡(a))(b-a)/(\log(b)-\log(a)) and is used to make predictions.

To summarize, and going back to the setting with KK ensemble members and mm-dimensional λk{\boldsymbol{\lambda}}_{k}’s, the optimization of {ξk,t}k=1K\{{\boldsymbol{\xi}}_{k,t}\}_{k=1}^{K} in the validation step involves 2mK2mK parameters, i.e., the lower/upper bounds for each hyperparameter and for each ensemble member (in practice, K≈5K\approx 5 and m≈5−10m\approx 5-10).

Appendix C Further details about the MLP and LeNet experiments

We provide in this section additional material about the experiments based on MLP and LeNet.

MLP: The multi-layer perceptron is composed of 2 hidden layers with 200 units each. The activation function is ReLU. Moreover a dropout layer is added before the last layer.

LeNet : This convolutional neural network is composed of a first conv2D layer (32 filters) with a max-pooling operation followed by a second conv2D layer (64 filters) with a max-pooling operation and finally followed by two dense layers (512 and number-of-classes units). The activation function is ReLU everywhere. Moreover, we add a dropout layer before the last dense layer.

As briefly discussed in the main paper, in the first tuning setting (i), there are two L2L_{2} regularization parameters for those models: one for all the weight matrices and one for all the bias terms of the conv2D/dense layers; in the second tuning setting (ii), the L2L_{2} regularization parameters are further split on a per-layer basis (i.e., a total of 3×2=63\times 2=6 and 4×2=84\times 2=8 L2L_{2} regularization parameters for MLP and LeNet respectively).

The ranges for the dropout and L2L_{2} parameters are [10−3,0.9][10^{-3},0.9] and [10−3,103][10^{-3},10^{3}] across all settings (i)-(ii), models and datasets (CIFAR-100 and Fashion MNIST).

We take the official train/test splits of the two datasets, and we further subdivide (80%/20%) the train split into actual train/validation sets. We use everywhere Adam with learning rate 10−410^{-4}, a batchsize of 256 and 200 (resp. 500) training epochs for LeNet (resp. MLP). We tune all methods to minimize the validation NLL. All the experiments are repeated with 3 random seeds.

C.2 Selection of the hyperparameters of batch ensemble

Following the recommendations from , we tuned

The type of the initialization of the vectors rk{\mathbf{r}}_{k}’s and sk{\mathbf{s}}_{k}’s (see Section 2.1). We indeed observed that the performance was sensitive to this choice. We selected from the different initialization schemes proposed in

Entries distributed according to the Gaussian distribution N(1,0.5×I)\mathcal{N}({\mathbf{1}},0.5\times{\mathbf{I}})

Entries distributed according to the Gaussian distribution N(1,0.75×I)\mathcal{N}({\mathbf{1}},0.75\times{\mathbf{I}})

Random independent signs, with probability of +1+1 equal to 0.50.5

Random independent signs, with probability of +1+1 equal to 0.750.75

A scale factor κ\kappa to make it possible to reduce the learning rate applied to the vectors rk{\mathbf{r}}_{k}’s and sk{\mathbf{s}}_{k}’s. Following , we considered the scale factor κ\kappa in {1.0,0.5}\{1.0,0.5\}.

Whether to use the Gibbs or ensemble cross-entropy at training time. Early experiments showed that Gibbs cross-entropy was substantially better so that we kept this choice fixed thereafter.

Whether to regularize the vectors rk{\mathbf{r}}_{k}’s and sk{\mathbf{s}}_{k}’s. mentioned that the two options perform equally well while we observed in those smaller-scale experiments that batch ensemble could overfit in absence of regularization.

The two batch ensemble-specific hyperparameters above (initialization type and κ\kappa) together with the MLP/LeNet hyperparameters were tuned by 50 trials of random search, separately for each ensemble size (3 and 5) and for each triplet (dataset, model type, tuning setting).

C.3 Selection of the hyperparameters of self-tuning networks

We re-used as much as possible the hyperparameters and design choices from , i.e., 5 warm-up epochs (during which no tuning happens) before starting the alternating scheme (2 training steps followed by 1 tuning step).

For the tuning step, the batch size is taken to be the same as that of the training step (256), while the learning was set to 5×10−45\times 10^{-4}.

We tuned the entropic regularization parameter τ∈{0.01,0.001,0.0001}\tau\in\{0.01,0.001,0.0001\}, separately for each triplet (dataset, model type, tuning setting), as done for all the methods compared in the benchmark. We observed that τ=0.001\tau=0.001 was often found to be the best option, and it therefore constitutes a good default value, as reported in .

As studied in Section C.5, we fix the embedding model e(⋅){\mathbf{e}}(\cdot) to be an MLP with one hidden layer of 64 units and a tanh activation.

C.4 Selection of the hyperparameters of hyper-batch ensemble

We followed the very same protocol as that used for the standard self-tuning network (as described in Section C.3).

By construction, we also inherit from the batch ensemble-specific hyperparameters (see Section C.2). To keep the protocol simple, we only tune the most important hyperparameter, namely the type of the initialization of the rank-1 terms (while the scale factor κ\kappa to discount the learning rate was not considered). As for any other methods in the benchmark, τ\tau and the initialization type were tuned separately for each triplet (dataset, model type, tuning setting).

For good default choices, we recommend to take τ=0.001\tau=0.001 and use an initialization scheme with random independent signs (with the probability of +1+1 equal to 0.750.75).

C.5 Choice of the embedding 𝐞​(⋅)𝐞⋅{\mathbf{e}}(\cdot)

We study the impact of the choice of the model that defines the embedding e(⋅){\mathbf{e}}(\cdot).

In , e(⋅){\mathbf{e}}(\cdot) is taken to be a simple linear transformation. In a slightly different context, the authors of consider MLPs with one hidden layer of 128 or 256 units, depending on their applications.

In the light of those previous choices, we compare the performance of different architectures of e(⋅){\mathbf{e}}(\cdot), namely linear (i.e., 0 units) and one hidden layer of 64, 128, and 256 units. The results are summarized in Figure 4-(left), for different ensemble sizes (one corresponding to the standard self-tuning networks ). We computed the validation NLL averaged over all the datasets (Fashion MNIST/CIFAR 100), model types (MLP/LeNet), tuning settings and random seeds.

Based on Figure 4-(left), we select for e(⋅){\mathbf{e}}(\cdot) an MLP with a single hidden layer of 64 units and a tanh activation function.

C.6 Sensitivity analysis with respect to the entropy regularization parameter τ𝜏\tau

We study the impact of the choice of the entropy regularization parameter τ\tau in (9). We report in Figure 4-(right) how the validation negative log-likelihood—aggregated over all the datasets (Fashion MNIST/CIFAR 100), model types (MLP/LeNet), tuning settings and random seeds—varies with τ∈{0.01,0.001,0.0001}\tau\in\{0.01,0.001,0.0001\}.

As discussed in Section C.3 and in Section C.2, a good default value, as already reported in is τ=0.001\tau=0.001.

C.7 Complementary results

In Table 3 and Table 4 (the latter table contains the efficient ensemble methods), we complete Table 9 with the addition of the results for the ensembles of size 5. To ease the comparison across different ensemble sizes, we incorporate as well the results for the size 3.

The conclusions highlighted in the main paper also hold for the larger ensembles of size 5. In Table 4, we can observe that hyper-batch​ ens with 5 members does not consistently improve upon its counterpart with 3 members. This trend is corrected if more training epochs are considered (see in Table 7 the effect of twice as many training epochs).

C.7.2 Assessment of the statistical significance of the results

To assess the statistical significance of the improvements displayed in Table 1, Table 3 and Table 4, we run the Wilcoxon signed-rank test, paired along settings, datasets and model types. We report the results in Table 5. The pairing of the tests is especially important for the comparisons between deep​ ens, fixed​ init​ hyper​ ens and hyper-deep​ ens since their respective performances are heavily conditioned on the initial random searches they build upon.

First, we can see that hyper-deep​ ens significantly improves upon both deep​ ens and fixed​ init​ hyper​ ens (with larger p-values in the latter case, though). Second, while hyper-batch​ ens significantly improves upon STN, hyper-batch​ ens can only be shown to be better than batch​ ens in terms of likelihood (with a 5% significance level). Overall, we also observe that we do not have significant improvements with respect to ECE which is known to be more noisy .

C.7.3 Diversity analysis

In this section, we study the diversity of the predictions made by the ensemble approaches from the experiments of Section 5.1.

To this end, we use the predictive disagreement metric from . This metric is based on the average of the pairwise comparisons of the predictions across the ensemble members. For a given pair of members, it is zero when they are making identical predictions, and one when all their predictions differ. We also normalize the diversity metric by the error rate (i.e., one minus the accuracy) to avoid the case where random predictions provide the best diversity.

For ensemble sizes 3 and 5, we compare in Table 6 the approaches hyper-deep ensemble, deep ensemble, hyper-batch ensemble and batch ensemble with respect to this metric. We can draw the following conclusions:

hyper-deep ensemble vs. deep ensemble: Compared to deep ensemble, we can observe that hyper-deep ensemble leads to significantly more diverse predictions, across all combination of (dataset, model type) and ensemble sizes. Moreover, we can also see that the diversity only slightly increases for deep ensemble going from 3 to 5 members, while it increases more markedly for hyper-deep ensemble. We hypothesise this is due to the more diverse set of models (with varied initialization and hyperparameters) that hyper-deep ensemble can tap into.

hyper-batch ensemble vs. batch ensemble: The first observation is that in this setting (the observation turns out to be different in the case of the Wide Resnet 28-10 experiments), batch ensemble leads to the largest diversity in predictions compared to all the other methods. Although lower compared with batch ensemble, the diversity of hyper-batch ensemble is typically higher than, or competitive with the diversity of deep ensembles.

C.7.4 Further comparison between batch ensemble and hyper-batch ensemble

As described in Section 4.1, the structure of the layers of hyper-batch​ ens leads to a 2x increase in memory compared with standard batch​ ens.

In an attempt to fairly account for this difference in memory footprints, we combine two batch ensemble models trained separately and whose total memory footprint amounts to that of hyper-batch​ ens. This procedure leads to ensembles with 6 and 10 members to compare to hyper-batch​ ens instantiated with 3 and 5 members respectively. To also normalize the training budget, hyper-batch​ ens is given twice as many training epochs as each of the batch​ ens models.

Table 7 presents the results of that comparison. In an nutshell, hyper-batch​ ens either continues to improve upon, or remain competitive with, batch​ ens, while still having the advantage of automatically tuning the hyperparameters of the underlying model (MLP or LeNet).

C.7.5 Ablation study about hyper-deep ensemble

In this section, we conduct two ablation studies about hyper-deep ensemble to better understand its components. We first focus on the effect of using the greedy algorithm of compared with the top-KK procedure used in . Second, we relate Algorithm 1 to the NES-RS procedure concurrently proposed by .

Starting from the set of models generated by random search (according to the setting of Section 5.1), we apply both the greedy and top-KK selection strategies, as previously used in , to form ensembles of size 5. We report the results of the evaluations of those strategies in Figure 5.

We can observe that the greedy procedure outperforms the top-KK procedure. While the former has an objective aware of the ensemble performance, the latter selects the models based only on their individual performance.

We still focus on the setting of Section 5.1, with ensembles of size 3 and 5. We study the value of the stratification step in Algorithm 1. To this end, we consider the following comparison that accounts for the total number of trained models:

hyper​ ens​ (70): Random search with 70 models followed by the greedy procedure of . Note that there is no stratification step in this variant. The resulting method falls back to NES-RS from where the architecture is kept fixed while hyperparameters are varied.

hyper-deep​ ens: The procedure described in Algorithm 1 that uses stratification and starts from 50 models obtained by random search (as used in the experiments of Section 5.1). Note that even though we need to stratify 5 models with 5 seeds, i.e., 525^{2}=25 models, we can reuse 5 models from the initial random search so that the total budget is 50+20=70 models to train (plus the cost of the calls to the greedy algorithm which is assumed negligible). The two approaches (A)-(B) therefore involve the same number of models to train.

The results of the comparison are reported in Table 8. While hyper-deep​ ens works slightly better, the differences with hyper​ ens​ (70) are not substantial. In the setting of Section 5.1, it thus appears that, provided that the initial random search produces enough models, the stratification step may be bypassed. In practice, this scheme, without stratification, can also be more convenient to implement.

C.7.6 Addendum to the results of Table 1

In Table 9, we complete the results of Table 1 with the addition of the Brier scores. Moreover, we provide the details of the performance of rand​ search and Bayes​ opt since only their aggregated best results were reported in Table 1.

Appendix D Further details about the ResNet experiments

We first explain the setting we used for training the ResNet 20 and Wide ResNet 28-10 architectures in Section 5 and conclude with the results of an empirical study over different algorithmic choices.

In the following we present the details for our training procedures. A similar training setup to ours for batch​ ens based on a Wide ResNet architecture can be found in the uncertainty-baselines repositoryhttps://github.com/google/uncertainty-baselines/tree/master/baselines/cifar.

For all methods (hyper-batch​ ens, batch​ ens, hyper-deep​ ens and deep​ ens), we optimize the model parameters using stochastic gradient descent (SGD) with Nesterov momentum of 0.90.9. For the ResNet 20 model we decay the learning rate by a factor of 0.1 after the epochs {80,180,200}\{80,180,200\} and for the Wide ResNet 28-10 model by a factor of 0.2 after the epochs {100,200,225}\{100,200,225\}. For tuning the hyperparameters in hyper-batch​ ens, we use Adam with a fixed learning rate. For hyper-batch​ ens, we use 95% of the data for training and the remaining 5% for optimizing the hyperparameters λ{\boldsymbol{\lambda}} in the tuning step. For the other methods we use the full training set.

For the efficient ensemble methods (hyper-batch​ ens and batch​ ens), we initialize the rank-1 factors, i.e., rksk⊤{\mathbf{r}}_{k}{\mathbf{s}}_{k}^{\top} and ukvk⊤{\mathbf{u}}_{k}{\mathbf{v}}_{k}^{\top} in (7), with entries independently sampled according to N(1,0.5)\mathcal{N}(1,0.5) for ResNet 20 and sampled according to N(1,1)\mathcal{N}(1,1) for Wide ResNet 28-10.

We make two minor adjustments of our model to adapt to the specific structure of the highly overparametrized ResNet models. First, we find that coupling the rank-1 factors corresponding to the hyperparamters to the rank-1 factors of weights is beneficial, i.e. we set uk:=rk{\mathbf{u}}_{k}:={\mathbf{r}}_{k} and vk:=sk{\mathbf{v}}_{k}:={\mathbf{s}}_{k}. This slightly decreases the flexibility of hyper-batch​ ens and makes it more robust against overfitting.

Second, we exclude the rank-1 factors from being regularized. In the original paper introducing batch​ ens , the authors mention that both options were found to work equally well and they finally choose not to regularize the rank-1 factors (to save extra computation). In our setting, we observe that this choice is important and regularizing the rank-1 factors leads to worse performance (a detailed analysis is given in Section D.2). Hence, we do not include the rank-1 factors in the regularization.

For hyper-batch​ ens we usually start with a log-uniform distribution over the hyperparameters ptp_{t} over the full range for the given bounds of the hyperparameters. For the ResNet models we find that reducing the initial ranges of ptp_{t} for the L2L_{2} regularization parameters by one order of magnitude is more stable (but we keep the original bounds for clipping the parameters).

We perform an exhaustive ablation of the different algorithmic choices for hyper-batch​ ens as well as for batch​ ens using the validation set. We run a grid search procedure evaluating up to five different values for each parameter listed below and repeat each run three times using different seeds. We find that the following configuration works best.

Initialization of each entry of the fast weights according to N(1,1)\mathcal{N}(1,1).

We multiply the learning rate for the fast weights by: 2.02.0.

Parameters specific to hyper-batch ensemble:

Range for the L2L_{2} parameters: [0.1,100][0.1,100].

Range for the label smoothing parameter: [0,0.2][0,0.2].

Entropy regularization parameter: τ=10−3\tau=10^{-3} (as also used in the other experiments and used by ).

Learning rate for the tuning step (where we use Adam): 10−510^{-5}.

Remarkably, we find that the shared set of parameters which work best for batch ensembles, also work best for hyper-batch ensembles. This makes our method an easy-to-tune drop-in replacement for batch ensembles.

D.2 Regularization of the rank-1 factors

As explained in the previous section, we find that for the Wide ResNet architecture, both hyper-batch ensemble and batch ensemble work best when the rank-1 factors (rksk⊤{\mathbf{r}}_{k}{\mathbf{s}}_{k}^{\top} and ukvk⊤{\mathbf{u}}_{k}{\mathbf{v}}_{k}^{\top}) are not regularized. We examine the performance of both models when using a regularization of the rank-1 factors. For these versions of the models, we run an ablation over the same algorithmic choices as done in the previous section. The results are displayed in Table 10. The performance of both methods is substantially worse than the unregularized versions as presented in the main text, Table 2.

D.3 Out-of-distribution experiments

In this section, we provide an out-of-distribution evaluation along the line of Table 1 in . More precisely, for each of the four approaches deep​ ens, hyper-deep​ ens, batch​ ens and hyper-batch​ ens, we compute on out-of-distribution samples from other image datasets the following metrics:

Mean maximum confidence (MMC) on out-distribution samples (lower is better)

The AUC of the ROC curve (AUROC) for the task of discriminating between in- and out-distributions based on the confidence value (higher is better)

The false positive rate at 95% true positive rate (FPR@95) in the same discriminative task (lower is better).

We summarize the results in Table 11, where we consider models both trained on CIFAR-10 (with evaluation on CIFAR-100 and SVHN) and CIFAR-100 (with evaluation on CIFAR-10 and SVHN). In a nutshell, hyper-deep​ ens (respectively hyper-batch​ ens) tends to favourably compare with deep​ ens (respectively batch​ ens) on CIFAR-10 and CIFAR-100, while they appear to perform worse over SVHN.

D.4 Complementary results for CIFAR-10

In this section we show complementary results to those presented in the main text for CIFAR-10. Figure 6 compares hyper-deep ensembles against deep ensembles for varying ensemble sizes. We find that the performance gain on CIFAR-10 is not as substantial as on CIFAR-100 presented in Figure 1. However, hyper-deep​ ens improves upon deep​ ens for large ensemble sizes in terms of NLL (cross entropy) and expected calibration error (ECE). The accuracy of hyper-deep​ ens is slightly higher for most ensemble sizes (except for ensemble sizes 3 and 10).

Figure 7 shows a comparison of additional metrics on the out of distribution experiment presented in the main text, Figure 3. We observe the same trend as in Figure 3 that hyper-batch​ ens is more robust than batch​ ens as it typically leads to smaller worst values (see top whiskers in the boxplot).

D.5 Complementary results for CIFAR-100

In this section we show complementary results to those presented in the main text for CIFAR-100. Figure 8 presents additional metrics (Brier score and expected calibration error) for varying ensemble sizes for hyper-deep ensemble and deep ensemble. Additionally to the strong improvements in terms of accuracy and NLL presented in Figure 1, we find that hyper-deep​ ens also improves in terms of Brier score and but is slightly less calibrated than deep ensemble for large ensemble sizes.

D.6 Memory and training time cost

For hyper-batch ensemble and batch ensemble, Table 12 reports the training time and memory cost in terms of number of parameters. Our method is roughly twice as costly as batch ensemble with respect to those two aspects. As demonstrated in the main text, this comes with the advantage of achieving better prediction performance. In Section C.7.4 we show that doubling the number of parameters for batch ensemble still leads to worse performance than our method.

Appendix E Towards more compact self-tuning layers

The goal of this section is to motivate the introduction of different, more compact parametrizations of the layers in self-tuning networks.

In , the choice of their parametrization (i.e., shifting and rescaling) is motivated by the example of ridge regression whose solution is viewed as a particular 2-layer linear network (see details in Section B.2 of ). The parametrization is however not justified for other losses beyond the square loss. Moreover, by construction, this parametrization leads to at least a 2x memory increase compared to using the corresponding standard layer.

If we take the example of the dense layer with input and output dimensions rr and ss respectively, recall that we have

Let us denote by ζj∈{0,1}s{\boldsymbol{\zeta}}_{j}\in\{0,1\}^{s} the one-hot vector with its jj-th entry equal to 1 and 0 elsewhere, ej(λ)e_{j}({\boldsymbol{\lambda}}) the jj-th entry of e(λ){\mathbf{e}}({\boldsymbol{\lambda}}) and δj{\boldsymbol{\delta}}_{j} the jj-th column of Δ{\boldsymbol{\Delta}}. We can rewrite the above equation as

As a result, we can re-interpret the parametrization of as a very specific linear combination of parameters Wj{\mathbf{W}}_{j} where the coefficients of the combination, i.e., e(λ){\mathbf{e}}({\boldsymbol{\lambda}}), depend on λ{\boldsymbol{\lambda}}.

Formulation (15) comes with two benefits. On the one hand, it reduces the memory footprint, as controlled by the low-rank factor hh which impacts the size of both (G∘e(λ))H⊤({\mathbf{G}}\circ{\mathbf{e}}({\boldsymbol{\lambda}})){\mathbf{H}}^{\top} and e(λ){\mathbf{e}}({\boldsymbol{\lambda}}). On the other hand, we can hope to get more expressiveness and flexibility since in (14), only the δj{\boldsymbol{\delta}}_{j}’s are learned, while in (15), both the vector gj{\mathbf{g}}_{j}’s and hj{\mathbf{h}}_{j}’s are learned.

Along the line of , but with a broader scope, beyond the ridge regression setting, we now provide theoretical arguments to justify the use of such a parametrization. We focus on the linear case with arbitrary convex loss functions. We start by recalling some notation, some of which slightly differ from the rest of the paper.

In the following derivations, we will use

The distribution over pair (x,y)({\mathbf{x}},y) is denoted by P\mathcal{P}

The distribution over hyperparameters (λ0,λ1)(\lambda_{0},{\boldsymbol{\lambda}}_{1}) is denoted by Q\mathcal{Q}

In a nutshell, we want to show that, for any λ∈Λ{\boldsymbol{\lambda}}\in\Lambda, Ue(λ){\mathbf{U}}{\mathbf{e}}({\boldsymbol{\lambda}})—i.e., a linear combination of parameters whose combination depends on λ{\boldsymbol{\lambda}}, as in (15)—can well approximate the solution w(λ){\mathbf{w}}({\boldsymbol{\lambda}}) of

In Proposition E.3, we show that when we apply a stochastic optimization algorithm to (16), e.g., SGD or variants thereof, with solution U^\hat{{\mathbf{U}}}, it holds in expectation over λ∼Q{\boldsymbol{\lambda}}\sim\mathcal{Q} that w(λ)≈U^e(λ){\mathbf{w}}({\boldsymbol{\lambda}})\approx\hat{{\mathbf{U}}}{\mathbf{e}}({\boldsymbol{\lambda}}) under some appropriate assumptions.

Our analysis operates with a fixed feature transformation ϕλ1\phi_{{\boldsymbol{\lambda}}_{1}} (e.g., a pre-trained network) and with a fixed embedding of the hyperparameters e{\mathbf{e}} (e.g., a polynomial expansion). In practice, those two quantities would however be learnt simultaneously during training. We stress that, despite those two technical limitations, the proposed analysis is more general than that of , in terms of both the loss functions and the hyperparameters covered (in , only the squared loss and λ0\lambda_{0} are considered).

We define (remembering the definition λ=(λ0,λ1){\boldsymbol{\lambda}}=(\lambda_{0},{\boldsymbol{\lambda}}_{1}))

E.2 Assumptions

For all λ∈Λ{\boldsymbol{\lambda}}\in\Lambda, gλ(⋅)g_{\boldsymbol{\lambda}}(\cdot) is convex and has LλL_{\boldsymbol{\lambda}}-Lipschitz continuous gradients.

E.3 Direct consequences

Under the assumptions above, we have the following properties:

For all λ∈Λ{\boldsymbol{\lambda}}\in\Lambda, the problem

admits a unique solution which we denote by w(λ){\mathbf{w}}({\boldsymbol{\lambda}}). Moreover, it holds that

F(⋅)F(\cdot) is strongly convex (C≻0{\mathbf{C}}\succ{\mathbf{0}}) and the problem

admits a unique solution which we denote by U⋆{\mathbf{U}}^{\star}.

E.4 Preliminary lemmas

which is the residual of the first-order Taylor expansion of gλ(⋅)g_{\boldsymbol{\lambda}}(\cdot) at w(λ){\mathbf{w}}({\boldsymbol{\lambda}}). Given Assumption (A1), it notably holds that

where in the last line we have used the optimality condition (17) of w(λ){\mathbf{w}}({\boldsymbol{\lambda}}).

As a direct consequence, we have the following result:

E.5 Main proposition

Before presenting the main result, we introduce a key quantity that will drive the quality of our guarantee. To measure how well we can approximate the family of solutions {w(λ)}λ∈Λ\{{\mathbf{w}}({\boldsymbol{\lambda}})\}_{{\boldsymbol{\lambda}}\in\Lambda} via the choice of e{\mathbf{e}} and Q\mathcal{Q}, we define

The definition is unique since according to (A2), we have Σ≻0{\boldsymbol{\Sigma}}\succ{\mathbf{0}}.

Let assume we have an, possibly stochastic, algorithm A\mathcal{A} such that after tt steps of A\mathcal{A} to optimize (16), we obtain Ut{\mathbf{U}}_{t} satisfying

for some tolerance εtA≥0\varepsilon_{t}^{\mathcal{A}}\geq 0 depending on both tt and the algorithm A\mathcal{A}. Denoting by Δt(λ)=Ute(λ)−w(λ){\boldsymbol{\Delta}}_{t}({\boldsymbol{\lambda}})={\mathbf{U}}_{t}{\mathbf{e}}({\boldsymbol{\lambda}})-{\mathbf{w}}({\boldsymbol{\lambda}}) the gap between the estimated and actual solution w(λ){\mathbf{w}}({\boldsymbol{\lambda}}) for any λ∈Λ{\boldsymbol{\lambda}}\in\Lambda, it holds that

Similarly, by definition of U⋆{\mathbf{U}}^{\star} as the minimum of F(⋅)F(\cdot), we have

Chaining the two inequalities leads to the first conclusion. The second conclusion stems from the application of (18). ∎