Squared Earth Mover's Distance-based Loss for Training Deep Neural Networks
Le Hou, Chen-Ping Yu, Dimitris Samaras
Introduction
Deep neural networks (DNNs) have become the preferred method for most machine learning applications, due to their ability to automatically learn optimal features and classifiers from inputs in an end-to-end fashion. In addition to superior performance as compared to conventional approaches, another reason for the popularity is their wide-range of applicability which includes convolutional neural networks (CNNs) for computer vision , recurrent neural networks (RNNs) for natural language processing , hybrid networks that combine CNN and RNN layers for speech recognition and audio processing , and more. In general, most DNNs are trained under one of two tasks: regression and classification. In a regression task, the network learns to generate a real-valued output that matches the ground-truth . In a classification task, the network learns to categorize an input to one of the training classes . Other tasks such as detection and segmentation are often cast as classification tasks using a sliding window or a multi-label classification approach.
To train a multi-class single-label classification network, softmax cross-entropy loss is by far the most popular loss function for the training regime, where the ground-truth is a binary vector consisting of a value 1 at the correct class index, and 0s everywhere else . During training, the objective is to minimize the negative log-likelihood of the loss by multiplying the network’s predictions to the binary ground-truth vectors. While the softmax cross-entropy loss has been successfully used across all applicable fields, the loss function does not take into account inter-class relationships which can be very informative. For example (Fig. 1), we want to estimate human age-groups from face images. Given an image of an adult, if network A’s output probability of the adult class is 0.1 with a max probability of 0.2 at the baby class, and network B’s output probability of the adult class is 0.1 with a max probability of 0.2 at the teen class, both networks would achieve the same softmax cross-entropy loss, while network B’s output probabilities are clearly closer to the ground-truth.
In this work, we show how the exact squared Earth Mover’s Distance (EMD) can be applied both as a stand-alone loss function or as a regularization term for multi-class classification problems using CNNs. The EMD is also known as the Wasserstein distance , which is the minimal cost required to transform one distribution to another . Recent work formulated an approximate Wasserstein loss for supervised multi-class multi-label learning using a linear model, and applied it to classification problems with predefined inter-class similarity metrics . In contrast, we show that an exact (without approximation) squared EMD (EMD2)-based loss exists for training single-label deep learning models directly, either with known inter-class relationship or without any priors on the inter-class relationships. We choose to use EMD2 instead of EMD as the loss function because squaring usually leads to faster convergence with gradient descent . Our experiments show that CNNs trained with our EMD2 loss result in better performance than CNNs with the standard softmax cross-entropy loss, and achieve state-of-the-art results on multiple datasets.
Computing the EMD requires a predefined ground distance matrix that quantifies the dissimilarities between classes. Existing work assumes a ground distance matrix. For example, when classes are ordered, the ground distance matrix has one dimensional embedding. Therefore, a closed-form solution exists for computing the EMD . Otherwise, the ground distance matrix was obtained externally . These approaches are limited by the need for reasonable assumptions about the ground distance matrix.
In this work, we also show how to learn the ground distance matrix using the CNN’s own features during training with limited additional computational cost, obviating the need for initial assumptions about the inter-class relationships. We propose to use EMD2 computed using the estimated ground distance matrix as a regularization term in CNN training. We call this self-guided training with EMD2-based regularization. We verify the learned matrix on datasets with known ordered-classes. Examples of such datasets include age estimation , image aesthetics , facial attractiveness prediction and others . Experiments show that our self-guided CNN with EMD2-based regularization performs as well as CNNs with EMD2-based loss computed using ground distance matrices based on prior knowledge. Furthermore, we show that on datasets with weak inter-class relationships, the learned ground distance matrix does not capture spurious inter-class relationships that could adversely affect performance.
We claim three major contributions in this paper:
For the first time, we show how an exact EMD2-based loss function can be used to train CNNs.
For the first time, we propose a method to efficiently discover inter-class relationships during training and use the discovered inter-class relationships as a ground distance matrix for CNN training with an exact EMD2-based regularization.
We improve state-of-the-art performance on datasets with strong inter-class relationships and avoid adverse effects on datasets with weak inter-class relationships.
EMD2-based Loss
In this section, we first introduce the standard softmax cross-entropy loss and discuss its drawbacks in detail. Then, we formulate the EMD2 and show it can be used as a loss function for classification problems with known assumptions on inter-class relationships.
It is obvious to see that the backpropagation of a DNN with cross-entropy loss only depends on . This is less robust compared to a loss function that depends on all entries of as argued in Fig. 1
2 EMD2-based loss on ordered-classes
Here, we first define the Earth Mover’s Distance (EMD), and explain how an EMD2-based loss function models inter-class relationships. Then, we define the problem of ordered-class classification and show when the exact EMD2 function can be computed by a closed-form equation.
We assume that a well performing CNN should predict class distributions such that classes closer to the ground truth class should have higher predicted probabilities than classes that are further away. We formulate this using the Earth Mover’s Distance (EMD). The EMD is defined as the minimum cost to transport the mass of one distribution (histogram) to the other.
Mass transportation defines the problem of transporting mass from a set of supplier clusters to a set of consumer clusters. Its formal definition is: Let be the supplier signature (distribution or histogram) with clusters (bins), where represents each cluster and is the mass (value) in each cluster. Let be the consumer signature. Let be the ground distance matrix where its -th entry is the distance between and . Matrix is usually defined as the -norm distance between clusters:
Let be the transportation matrix where its -th entry indicates the mass transported from to . A valid transportation satisfies four constraints. First, the amount of mass transported must be positive. Second, the amount of mass transported from a supplier cluster must not exceed its total mass. Third, the amount of mass transported to a consumer cluster must not exceed its total mass. Finally, the total flow must not exceed the total mass that can be transported. These four conditions can be summarized respectively below:
Under the constraints defined above, the overall cost of flow is defined as:
2.2 Ground distance matrix of ordered-classes
Computing the EMD between two distributions requires a predefined matrix, the ground distance matrix which is unknown in most cases. However, in classification tasks with ordered classes we can define . By ordered classes we mean classes that can be represented as real numbers, for example, human age ranges or aesthetic preference levels. The difference between ordered-class classification and regression is that in the problem of ordered-class classification, the ground truth labels and predictions are discrete. Hence, in practice better performance can be achieved using a multi-class classification model instead of a regression model, and the ground distances between those classes can be based on their inherent ordering.
Without loss of generality, we assume that in all ordered-class classification problems, the classes are ranked as and the distance between and is .
2.3 EMD2 loss for ordered-class classification
EMD has been shown to be equivalent to Mallows distance which has a closed-form solution , if the ground distance matrix and distributions and satisfy certain conditions, as shown in . We will show that these required conditions are satisfied in ordered-class classification problems.
The first condition is that the two distributions and to be compared must have equal mass: . Note that this condition is always satisfied if is produced by a softmax layer, as the output vector of a softmax layer is a normalized probability density function that sums to 1. And since the number of classes in the predicted distribution is the same as the target distribution, then .
The second condition is that the ground distance matrix must have an one-dimensional embedding. Assuming and are sorted according to their inherent rank values without loss of generality, this condition can be expresses as , for a constant and all , that . Clearly, this assumption can always be satisfied in ordered-class classification problems.
The third and final condition is that the distributions to be compared must be sorted vectors. This condition is also always satisfied since we assumed and are sorted without loss of generality. Then, based on the conclusion by Levina et al. , the normalized EMD can be computed exactly and in closed-form:
2.4 Derivative of the EMD2 loss with ordered-classes
The coefficients of can be propagated using the standard backpropagation method.
Comparing Eq. LABEL:eq:emd-dev-ordered with Eq. 1, we can see that the backpropagation of a network trained with cross-entropy loss is only based on and , whereas the backpropagation of a network trained with EMD2 loss is based on all elements of and .
Self-Guided EMD
Sec. 2.2.3 showed the formulation of EMD2 loss for ordered classes for which the ground distance matrix can be easily assumed. However, in general the matrix is unknown. In this section, we show how to compute a ground distance using empirical evidence. Furthermore, we propose a method that computes the EMD2-based loss between prediction and ground truth with time complexity for single-label classification problems. Finally, for classification problems in general, we propose to use EMD2 as a regularization term to the cross-entropy loss. Note that the proposed regularization term does not compete with other regularization terms such as weight decay. One can apply all of these regularization terms for training.
We focus on estimating the ground distance matrix expressed in Eq. 2. Note that for classification problems, the predicted classes (supplier signatures) and the ground truth classes (consumer signatures) are the same set of classes. In other words , , . We will use to indicate the same set of classes in the future. We estimate first by estimating all (equivalently ) directly. Then we compute an initial estimation of , denoted by directly from Eq. 2. Finally we postprocess to obtain an estimated .
To estimate each , we extract features on all instances of the -th class, and use the centroid of these feature vectors as an initial estimate of , denoted as . To extract the features of one instance, we follow a standard method : we use the second-to-last layer neural responses of the CNN that is being trained, as feature vectors. Note that we normalize each instance’s CNN features. Intuitively, because CNNs learn to linearly separate classes with the second-to-last layer features (there is no subsequent non-linearity), it is meaningful to average the feature vectors to compute the class centroids.
We denote as the initial estimation of the ground distance matrix where its -th entry . We observe in practice that the CNN features cannot provide sufficient class separation before the network has partially converged. As a result, many entries of matrix are close to zero, indicating that class centroids are not well separated. To address this, we map each row of onto uniformly distributed values: each entry is mapped to its percentile value in its row. Formally, denoting the transformed matrix as , this operation is formulated as:
Note that based on this definition, for all .
2 Self-guided EMD2 regularization
We now describe how to calculate the EMD between and , defined by Eq. 8. In the case of single-label classification, the consumer’s mass vector (target distribution) is a binary vector where only the index of the ground truth class equals to 1: . According to the constraint defined by Eq. 5, all mass must be transferred to the -th cluster. Thus, the transportation matrix must satisfy if , otherwise . The resulting EMD is:
The computational complexity of EMD by Eq. 14 is .
The exact EMD defined by Eq. 14 can be directly used as a loss function. Its derivative with respect to network parameters is:
In practice, the optimization does not converge to a desired local optimum using Eq. 14 directly. We observed that using Eq. 15 for gradient descent ends up lowering the predicted probabilities of all classes, which leads the model to converge to a local minimum with uniformly distributed predictions. To address this optimization problem, we modify Eq. 14 and use it as a regularizer instead of a stand-alone loss function. Additionally, for faster optimization with gradient descent, we use instead of as the mass for each supplier cluster. Our hybrid loss with EMD2-based regularization for classification problems is then defined as:
where , and are predefined parameters, such that defines the weight of EMD2 regularization, the power term determines the sensitivity of the ground distance: a very high means that EMD only penalizes predictions on classes that are far away from the ground truth class, and is the ground distance bias. In our experiments, we use a negative so that is negative, which means that the network will be rewarded for predictions that are closer to the ground truth class. Note that in Eq. 16, we omitted the regularization (weight decay) term which was used in experiments.
Experiments and Results
In this section, we show implementation details and experimental results. First, we demonstrate that our EMD2-based losses outperform the softmax cross-entropy loss on datasets with known strong inter-class relationships: datasets with ordered-classes. At the same time, we show that our self-guided EMD2-based regularization which does not require known inter-class relationships performs as well as the EMD2-based loss that requires strong assumptions on the inter-class relationships. Finally, we show that on datasets without strong inter-class relationships such as ImageNet , our learned ground distance matrix does not capture spurious inter-class relationships that could lower performance.
We test the EMD2-based losses on different network architectures including the AlexNet , VGG 16-layer network , and wide residual network . For optimization, we use stochastic gradient descent with momentum 0.98 in all experiments. The learning rates were selected from individually for each method on each dataset. We notice that when using EMD2 as a regularizer, predicted probability for some classes can be very close to zero, resulting in errors when computing the logarithm of the prediction vector. To solve this, we simply add to the predicted probabilities of all classes. For experiments on ImageNet, we use the same data augmentation and weight decay methods used by AlexNet . For experiments on all other datasets, we use the following data augmentation methods: during training, we first crop smaller images from the original images (translation augmentation); second, we perturb the images’ RGB colors slightly; third, the images are randomly flipped horizontally; fourth, we rotate the images by degrees; finally, we adjust the aspect ratio by . During testing, we use the average prediction from the center crop and its mirrored image. We use Theano for network implementation.
2 Computational complexity of EMD
Our implementation of EMD2-based loss functions adds less than 10% of CNNs’ training time for each iteration and no additional test time. The increment of training time is due to two introduced procedures. However, both of them add very marginal computational complexity. First, the computation time of the loss function is very limited compared to the training time of the entire CNN. Second, to compute the ground distance matrix, no additional CNN forwarding process is needed: we simply store the features of each instance during each training iteration. Moreover, as shown in Fig. 2, we find that CNNs with the EMD2-based losses achieves the same performance as CNNs with the cross-entropy loss with only of its training epochs.
3 Methods tested
We train the AlexNet (ALX) from scratch, to compare with the published baselines . For experiments on the Adience dataset, we test a smaller version of AlexNet following the baseline method . We name it as ALXs.
We fine-tune the VGG 16-layer network (VGG) pre-trained on ImageNet, to compare with the published baseline .
We use a 40-layer residual network with identity mapping and bottleneck design that is the same as . We train this network from scratch.
We fine-tune a RES pre-trained on ImageNet.
The tested loss functions are described below:
The regression loss. To use this loss function, the output neurons of a regression network has linear activation functions, instead of softmax, following the conventional regression CNN approach . Other parts of the regression network is identical to a classification network.
The EMD2 loss defined by Eq. 10 on ordered-class classification problems.
The self-guided EMD2 regularization training defined by Eq. 16. For the predefined parameters in Eq. 16, we choose and . We “jump-started” the networks by training with softmax cross-entropy for the first 4 epochs with , as a way to avoid using inaccurately estimated ground distance matrix for computing class representations. After 4 epochs we choose a such that the EMD2 term is 3 to 4 times smaller than the cross-entropy term. We find that in ordered-class datasets, the performance was not sensitive when we changed .
The self-guided EMD2 regularization training defined by Eq. 16, but with and and keep other parameters unchanged.
The approximate EMD loss . It requires a predefined ground distance matrix . Given a specific network, we used the final estimated by XEMD2 with the same network, as the predefined matrix. The number of matrix scaling iterations in is set to 100. The entropic regularizer in is selected from based on the validation error. We use a Caffe implementation of this loss function .
We test different loss functions on various networks. For example, ALX-XE is the AlexNet with the softmax cross-entropy loss.
4 Age estimation on Adience dataset
Age estimation using human face images is important for analyzing and understanding human faces . We test our method on the Adience dataset . The Adience dataset contains 26,000 images in 8 age-groups, and a five-fold cross-validation evaluation scheme. For comparison, we use a smaller version of AlexNet and fine-tune a pre-trained VGG 16-layer network as described in the earlier experimental details. We use the conventional accuracy of exact match (AEM%) and with-in-one-category-off match (AEO%) as evaluation metrics.
The results are shown in Tab. 1. The pre-trained VGG network fine-tuned using our proposed EMD2-based loss function (VGGF-XEMD2) outperforms the state-of-the-art Deep Expectation (VGGF-DEX) method on the same dataset. Our self-guided methods (XEMD1, XEMD2) perform as well as the method with prior knowledge (EMD). The VGGF-DEX + IMDB-WIKI method achieves better results using an external age estimation dataset IMDB-WIKI which is 10 times larger than Adience. Therefore, our method improves the state-of-the-art when training without external face datasets. The regression loss based methods have low AEMs. We believe the reason is that the loss is sensitive to outliers. We show the AEM and AEO results with respect to the number of training epochs in Fig. 2, and the learned ground distance matrix (Eq. 13) by our best performing method RESF-XEMD2 on the Adience dataset in Tab. 2. We can see that age-groups that are further away from each other always have larger ground distances than age-groups that are closer to each other. In addition, the ground distance matrix reflects similarities between different age-groups, e.g., people in their 15-20s are more similar to 25-32s than to 8-13s.
5 Age estimation on Images of Groups dataset
We further test our method on the Images of Groups dataset . This dataset contains 3,500 training face images and 1,000 testing face images in 7 age-groups. We train wide residual networks from scratch and fine-tun a pre-trained VGG 16-layer network on this dataset. The results are shown in Tab. 3. Our EMD2-based losses outperforms the cross-entropy loss and the loss (regression) in a majority of the cases, with our results achieving a new state-of-the-art on this dataset. Our self-guided methods (XEMD1, XEMD2) perform as well as the method with prior knowledge (EMD).
6 Image aesthetics
Assessing image aesthetics automatically has a wide range of applications . We test our method on the Image Aesthetics with Attributes Database (AADB) which contains 8,458 training and 1,000 testing images, labeled as real numbers in according to the viewers’ aesthetic judgments. To transform this dataset into a classification dataset, we discretize the real number labels to 10 bins, balancing the number of training images in each bin. During testing, we compute the expected aesthetic scores according to the predicted distributions of 10 aesthetic bins. This give us real-numbered predictions. We use Spearmans’ rank correlation as the evaluation metric, following .
The results are shown in Tab. 4. Our EMD2-based losses again outperform cross-entropy loss and loss (regression) significantly. We conduct additional experiments by discretizing the real-numbered aesthetic labels to 8 different number of bins (3,4,5,6,7,8,9,10 bins), which give us 8 sets of ground truth labels. Then, we fine-tune one VGGF-EMD network for each set of ground truth and average the prediction results into an ensemble model, and denote this method as VGGF-EMD 8. It achieves state-of-the-art results with only image training data, outperforming the previous state-of-the-art method trained with additional 11 labels such as color harmony and vivid color information.
7 Generalization on ImageNet
We show the generalization ability of our self-guided EMD2-based regularizer (Eq. 16) on the ImageNet ILSVRC 2012 dataset , which is a classification dataset with weak inter-class relationships. We test the original AlexNet and 40-layer residual network on this dataset with cross-entropy, and separately with the self-guided EMD2 regularization. We do not test the VGG network because training it from scratch is time consuming. The results of the validations set are reported in Tab. 5. We see that one can safely apply our self-guided method on datasets with weak inter-class relationships.
Our method does not outperform the baseline because our self-guided method finds that the inter-class relationship on this dataset is less significant compared to those in ordered-class datasets. We show this by measuring the Standard Deviation of pair-wise Distances between the estimated centroids of the classes (SDD, standard deviation between all entries of introduced in Sec. 3.1). A higher SDD indicates a stronger inter-class relationship. On the Adience dataset, SDD; On Images of Groups SDD; On AADB, SDD; On ImageNet, SDD.
Conclusion
In this work, we argued that the conventional softmax cross-entropy loss for training CNNs only maximizes the predicted probability at the ground truth label, and ignores the inter-class relationships. We proposed to use the exact squared earth mover’s distance (EMD2) in loss functions for CNN training to take class relationships into account. We evaluated our methods on two age estimation datasets and one image aesthetic assessment dataset. Our method significantly outperformed state-of-the-art regression-based and cross-entropy-based CNNs using no external datasets, and with only image information. Furthermore, we demonstrated that our method can discover the inter-class relationships efficiently with no prior knowledge. Finally, we showed that our method can be applied to datasets with weak inter-class relationships with no adverse results. Our future works include a more sophisticated ground distance matrix computing method, and exploring variants of EMDs that can perform better on datasets with weak inter-class relationships.