Stochastic Multiple Choice Learning for Training Diverse Deep Ensembles
Stefan Lee, Senthil Purushwalkam, Michael Cogswell, Viresh Ranjan, David Crandall, Dhruv Batra
Introduction
Perception problems rarely exist in a vacuum. Typically, problems in Computer Vision, Natural Language Processing, and other AI subfields are embedded in larger applications and contexts. For instance, the task of recognizing and segmenting objects in an image (semantic segmentation ) might be embedded in an autonomous vehicle , while the task of describing an image with a sentence (image captioning ) might be part of a system to assist visually-impaired users .
In this work, we fix the form of this mapping to be the union of outputs from an ensemble of predictors such that , and address the task of training ensemble members such that minimizes oracle loss. Under our formulation, different ensemble members are free to specialize on subsets of the data distribution, so that collectively they produce a set of outputs which covers the space of high probability predictions well.
Diverse solution sets are especially useful for structured prediction problems with multiple reasonable interpretations, only one of which is correct. Situations that often arise in practical systems include:
Implicit class confusion. The label space of many classification problems is often an arbitrary quantization of a continuous space. For example, a vision system may be expected to classify between tables and desks, despite many real-world objects arguably belonging to both classes. By making multiple predictions, this implicit confusion can be viewed explicitly in system outputs.
Ambiguous evidence. Often there is simply not enough information to make a definitive prediction. For example, even a human expert may not be able to identify a fine-grained class (e.g., particular breed of dog) given an occluded or distant view, but they likely can produce a small set of reasonable guesses. In such cases, the task of producing a diverse set of possibilities is more clearly defined than producing one correct answer.
Bias towards the mode. Many models have a tendency to exhibit mode-seeking behaviors as a way to reduce expected loss over a dataset (e.g., a conversation model frequently producing ‘I don’t know’). By making multiple predictions, a system can improve coverage of lower density areas of the solution space, without sacrificing performance on the majority of examples.
In other words, by optimizing for the oracle loss, a multiple-prediction learner can respond to ambiguity much like a human does, by making multiple guesses that capture multi-modal beliefs. In contrast, a single-prediction learner is forced to produce a solution with low expected loss in the face of ambiguity. Figure 1 illustrates how this can produce solutions that are not useful in practice. In semantic segmentation, for example, this problem often causes objects to be predicted as a mixture of multiple classes (like the horse-cow shown in the figure). In image captioning, minimizing expected loss encourages generic sentences that are ‘safe’ with respect to expected error but not very informative. For example, Figure 1 shows two pairs of images each having different image content but very similar, generic captions – the model knows it is safe to assume that birds are on branches and that cakes are eaten with forks.
In this paper, we generalize the Multiple Choice Learning paradigm to jointly learn ensembles of deep networks that minimize the oracle loss directly. We are the first to adapt these ideas to deep networks and we present a novel training algorithm that avoids costly retraining and learning difficulty of past methods. Our primary technical contribution is the formulation of a stochastic block gradient descent optimization approach well-suited to minimizing the oracle loss in ensembles of deep networks, which we call Stochastic Multiple Choice Learning (sMCL). Our formulation is applicable to any model trained with stochastic gradient descent, is agnostic to the form of the task dependent loss, is parameter-free, and is time efficient, training all ensemble members concurrently.
We demonstrate the broad applicability and efficacy of sMCL for training diverse deep ensembles with interpretable emergent expertise on a wide range of problem domains and network architectures, including Convolutional Neural Network (CNN) ensembles for image classification , Fully-Convolutional Network (FCN) ensembles for semantic segmentation , and combined CNN and Recurrent Neural Network (RNN) ensembles for image captioning . We provide detailed analysis of the training and output behaviors of the resulting ensembles, demonstrating how ensemble member specialization and expertise emerge automatically when trained using sMCL. Our method outperforms existing baselines and produces sets of outputs with high oracle performance.
Related Work
Ensemble Learning. Much of the existing work on training ensembles focuses on diversity between member models as a means to improve performance by decreasing error correlation. This is often accomplished by resampling existing training data for each member model or by producing artificial data that encourages new models to be decorrelated with the existing ensemble . Other approaches train or combine ensemble members under a joint loss . More recently, work of Hinton et al. and Ahmed et al. explores using ‘generalist’ network performance statistics to inform the design of ensemble-of-expert architectures for classification. In contrast, sMCL discovers specialization as a consequence of minimizing oracle loss. Importantly, most existing methods do not generalize to structured output labels, while sMCL seamlessly adapts, discovering different task-dependent specializations automatically.
Generating Multiple Solutions. There is a large body of work on the topic of extracting multiple diverse solutions from a single model ; however, these approaches are designed for probabilistic structured-output models and are not directly applicable to general deep architectures.
Most related to our approach is the work of Guzman-Rivera et al. which explicitly minimizes oracle loss over the outputs of an ensemble, formalizing this setting as the Multiple Choice Learning (MCL) paradigm. They introduce a general alternating block coordinate descent training approach which requires retraining models multiple times. More recently, Dey et al. reformulated this problem as a submodular optimization task in which ensemble members are learned sequentially in a boosting-like manner to maximize marginal gain in oracle performance. Both these methods require either costly retraining or sequential training, making them poorly suited to modern deep architectures that can take weeks to train. To address this serious shortcoming and to provide the first practical algorithm for training diverse deep ensembles, we introduce a stochastic gradient descent (SGD) based algorithm to train ensemble members concurrently.
Multiple-Choice Learning as Stochastic Block Gradient Descent
We consider the task of training an ensemble of differentiable learners that together produce a set of solutions with minimal loss with respect to an oracle that selects only the lowest-error prediction.
Minimizing Oracle Loss with Multiple Choice Learning. In order to directly minimize the oracle loss for an ensemble of learners, Guzman-Rivera et al. present an objective which forms a (potentially tight) upper-bound. This objective replaces the in the oracle loss with indicator variables where is 1 if predictor has the lowest error on example ,
The resulting minimization is a constrained joint optimization over ensemble parameters and data-point assignments. The authors propose an alternating block algorithm, shown in Algorithm 1, to approximately minimize this objective. Similar to K-Means or ‘hard-EM,’ this approach alternates between assigning examples to their min-loss predictors and training models to convergence on the partition of examples assigned to them. Note that this approach is not feasible with training deep networks, since modern architectures can take weeks or months to train a single model once.
Stochastic Multiple Choice Learning. To overcome this shortcoming, we propose a stochastic algorithm for differentiable learners which interleaves the assignment step with batch updates in stochastic gradient descent. Consider the partial derivative of the objective in Eq. 1 with respect to the output of the individual learner on example ,
Notice that if is the minimum error predictor for example , then , and the gradient term is the same as if training a single model; otherwise, the gradient is zero. This behavior lends itself to a straightforward optimization strategy for learners trained by SGD based solvers. For each batch, we pass the examples through the learners, calculating losses from each ensemble member for each example. During the backward pass, the gradient of the loss for each example is backpropagated only to the lowest error predictor on that example (with ties broken arbitrarily).
Experiments
In this section, we present results for sMCL ensembles trained for the tasks and deep architectures shown in Figure 3. These include CNN ensembles for image classification, FCN ensembles for semantic segmentation, and a CNN+RNN ensembles for image caption generation.
Baselines. Many existing general techniques for inducing diversity are not directly applicable to deep networks. We compare our proposed method against:
Classical ensembles in which each model is trained under an independent loss with differing random initializations. We will refer to these as Indp. ensembles in figures.
MCL that alternates between training models to convergence on assigned examples and allocating examples to their lowest error model. We repeat this process for 5 meta-iterations and initialize ensembles with (different) random weights. We find MCL performs similarly to sMCL on small classification tasks; however, MCL performance drops substantially on segmentation and captioning tasks. Unlike sMCL which can effectively reassign an example once per epoch, MCL only does this after convergence, limiting its capacity to specialize compared to sMCL. We also note that sMCL is 5x faster than MCL, where the factor 5 is the result of choosing 5 meta-iterations (other applications may require more, further increasing the gap.)
Dey et al. train models sequentially in a boosting-like fashion, each time reweighting examples to maximize marginal increase of the evaluation metric. We find these models saturate quickly as the ensemble size grows. As performance increases, the marginal gain and therefore the weights approach zero. With low weights, the average gradient backpropagated for stochastic learners drops substantially, reducing the rate and effectiveness of learning without careful tuning. To compute weights, requires an error measure bounded above by 1: accuracy (for classification) and IoU (for segmentation) satisfy this; the CIDEr-D score divided by 10 guarantees this for captioning.
Oracle Evaluation. We present results as oracle versions of the task-dependent performance metrics. These oracle metrics report the highest score over all outputs for a given input. For example, in classification tasks, oracle accuracy is exactly the top- criteria of ImageNet , i.e. whether at least one of the outputs is the correct label. Likewise, the oracle intersection over union (IoU) is the highest IoU between the ground truth segmentation and any one of the outputs. Oracle metrics allow the evaluation of multiple-prediction systems separately from downstream re-ranking or selection systems, and have been extensively used in previous work .
Our experiments convincingly demonstrate the broad applicability and efficacy of sMCL for training diverse deep ensembles. In all three experiments, sMCL significantly outperforms classical ensembles, Dey et al. (typical improvements of 6-10%), and MCL (while providing a 5x speedup over MCL). Our analysis shows that the exact same algorithm (sMCL) leads to the automatic emergence of different interpretable notions of specializations among ensemble members.
Model. We begin our experiments with sMCL on the CIFAR10 dataset using the small convolutional neural network “CIFAR10-Quick” provided with the Caffe deep learning framework . CIFAR10 is a ten way classification task with small 3232 images. For these experiments, the reference model is trained using a batch size of 350 for 5,000 iterations with a momentum of 0.9, weight decay of 0.004, and an initial learning rate of 0.001 which drops to 0.0001 after 4000 iterations.
Results. Oracle accuracy for sMCL and baseline ensembles of size 1 to 6 are shown in Figure 4(a). The sMCL trained ensembles result in higher oracle accuracy than the baseline methods, and are comparable to MCL while being 5x faster. The method of Dey et al. performs worse than independent ensembles as ensemble size grows. Figure 4(b) shows the oracle loss during training for sMCL and regular ensembles. The sMCL trained models optimize for the oracle cross-entropy loss directly, not only arriving at lower loss solutions but also reducing error more quickly.
Interpretable Expertise: sMCL Induces Label-Space Clustering. Figure 4(c) shows the class-wise distribution of the assignment of test datapoints to the oracle or ‘winning’ predictor for an sMCL ensemble. The level of class division is striking – most predictors become specialists for certain classes. Note that these divisions emerge from training under the oracle loss and are not hand-designed or pre-initialized in any way. In contrast, Figure 4(f) show that the oracle assignments for a standard ensemble are nearly uniform. To explore the space between these two extremes, we loosen the constraints of Eq. 1 such that the lowest error predictors are penalized. By varying between 1 and the number of ensemble members , the models transition from minimizing oracle loss at to a traditional ensemble at . Figures 4(d) and 4(e) show these results. We find a direct correlation between the degree of specialization and oracle accuracy, with netting highest oracle accuracy.
2 Semantic Segmentation
We now present our results for the semantic segmentation task on the Pascal VOC dataset .
Model. We use the fully convolutional network (FCN) architecture presented by Long et al. as our base model. Like , we train on the Pascal VOC 2011 training set augmented with extra segmentations provided in and we test on a subset of the VOC 2011 validation set. We initialize our sMCL models from a standard ensemble trained for 50 epochs at a learning rate of . The sMCL ensemble is then fine-tuned for another 15 epochs at a reduced learning rate of .
Results. Figure 5(a) shows oracle accuracy (class-averaged IoU) for all methods with ensemble sizes ranging from 1 to 6. Again, sMCL significantly outperforms all baselines (~ relative improvement over classical ensembles). In this more complex setting, we see the method of Dey et al. saturates more quickly – resulting in performance worse than classical ensembles as ensemble size grows. Though we expect MCL to achieve similar results as sMCL, retraining the MCL ensembles a sufficient number of times proved infeasible so results after five meta-iterations are shown.
Interpretable Expertise: sMCL as Segmentation Specialists. In Figure 5(b), we analyze the class distribution of the predictions using an sMCL ensemble with members. For each test sample, the oracle picks the prediction which corresponds to the ensemble member with the highest accuracy for that sample. We find the specialization with respect to classes is much less evident than in the classification experiments. As segmentation presents challenges other than simply selecting the correct class, specialization can occur in terms of shape and frequency of predicted segments in addition to class divisions; however, we do still see some class biases – network 2 captures cows, tables, and sofas well and network 4 has become an expert on sheep and horses.
Figure 6 shows qualitative results from a four member sMCL ensemble. We can clearly observe the diversity in the segmentations predicted by different members. In the first row, we see the majority of the ensemble members produce dining tables of various completeness in response to the visual uncertainty caused by the clutter. Networks 2 and 3 capture this ambiguity well, producing segmentations with the dining table completely present or absent. Row 2 demonstrates the capacity of sMCL ensembles to provide multiple high quality solutions. The models are confused whether the animal is a horse or a cow – models 1 and 3 produce typical ‘safe’ responses while models 2 and 4 attempt to give cohesive responses. Finally, row 3 shows how the models can learn biases about the frequency of segments with model 3 presenting only the sheep.
3 Image Captioning
In this section, we show that sMCL trained ensembles can produce sets of high quality and diverse sentences, which is essential to improving recall and capturing ambiguities in language and perception.
To summarize, we propose Stochastic Multiple Choice Learning (sMCL), an SGD-based technique for training diverse deep ensembles that follows a ‘winner-take-gradient’ training strategy. Our experiments demonstrate the broad applicability and efficacy of sMCL for training diverse deep ensembles. In all experimental settings, sMCL significantly outperforms classical ensembles and other strong baselines including the 5x slower MCL procedure. Our analysis shows that exactly the same algorithm (sMCL) automatically generates specializations among ensemble members along different task-specific dimensions. sMCL is simple to implement, agnostic to both architecture and loss function, parameter free, and simply involves introducing one new sMCL layer into existing ensemble architectures.
This work was supported in part by a National Science Foundation CAREER award, an Army Research Office YIP award, ICTAS Junior Faculty award, Office of Naval Research grant N00014-14-1-0679, Google Faculty Research award, AWS in Education Research grant, and NVIDIA GPU donation, all awarded to DB, and by an NSF CAREER award (IIS-1253549), the Intelligence Advanced Research Projects Activity (IARPA) via Air Force Research Laboratory contract FA8650-12-C-7212, a Google Faculty Research award, and an NVIDIA GPU donation, all awarded to DC. Computing resources used by this work are supported in part by NSF (ACI-0910812 and CNS-0521433), the Lily Endowment, Inc., and the Indiana METACyt Initiative. The U.S. Government is authorized to reproduce and distribute reprints for Governmental purposes notwithstanding any copyright annotation thereon. Disclaimer: The views and conclusions contained herein are those of the authors and should not be interpreted as necessarily representing the official policies or endorsements, either expressed or implied, of IARPA, AFRL, NSF, or the U.S. Government.