Geometric Dataset Distances via Optimal Transport
David Alvarez-Melis, Nicolò Fusi
Introduction
A key hallmark of machine learning practice is that labeled data from the application of interest is usually scarce. For this reason, there is vast interest in methods that can combine, adapt and transfer knowledge across datasets and domains. Entire research areas are devoted to these goals, such as domain adaptation, transfer-learning and meta-learning. A fundamental concept underlying all these paradigms is the notion of distance (or more generally, similarity) between datasets. For instance, transferring knowledge across similar domains should intuitively be easier than across distant ones. Likewise, given a choice of various datasets to pretrain a model on, it would seem natural to choose the one that is closest to the task of interest.
Despite its evident usefulness and apparent simpleness, the notion of distance between datasets is an elusive one, and quantifying it efficiently and in a principled manner remains largely an open problem. Doing so requires solving various challenges that commonly arise precisely in the settings for which this notion would be most useful, such as the ones mentioned above. For example, in supervised machine learning settings the datasets consist of both features and labels, and while defining a distance between the former is often —though not always— trivial, doing so for the labels is far from it, particularly if the label-sets across the two tasks are not identical (as is often the case for off-the-shelf pretrained models).
Current approaches to transfer learning that seek to quantify dataset similarity circumvent these challenges in various ingenious, albeit often heuristic, ways. A common approach is to compare the dataset via proxies, such as the learning curves of a pre-specified model (Leite & Brazdil, 2005) or its optimal parameters (Achille et al., 2019; Khodak et al., 2019) on a given task, or by making strong assumptions on the similarity or co-occurrence of labels across the two datasets (Tran et al., 2019). Most of these approaches lack guarantees, are highly dependent on the probe model used, and require training a model to completion (e. g., to find optimal parameters) on each dataset being compared. On the opposite side of the spectrum are principled notions of discrepancy between domains Ben-David et al. (2007); Mansour et al. (2009), which nevertheless are often not computable in practice, or do not scale to the type of datasets used in machine learning practice.
In this work, we seek to address some of these limitations by proposing an alternative notion of distance between datasets. At the heart of this approach is the use of optimal transport (OT) distances (Villani, 2008) to compare distributions over feature-label pairs in a geometrically-meaningful and principled way. In particular, we propose a hybrid Euclidean-Wasserstein distance between feature-label pairs across domains, where labels themselves are modeled as distributions over features vectors. As a consequence of this technique, our framework allows for comparison of datasets even if their label sets are completely unrelated or disjoint, as long as a distance between their features can be defined. This notion of distance between labels, a by-product of our approach, has itself various potential uses, e. g., to optimally sub-sample classes from large datasets for more efficient pretraining.
In summary, we make the following contributions:
We introduce a notion of distance between datasets that is principled, flexible and computable in practice
We propose various algorithmic shortcuts to scale up computation of this distance to very large datasets
We provide extensive empirical evidence that this distance is highly predictive of transfer learning success across various domains, tasks and data modalities
Related Work
Various notions of (dis)similarity between data distributions have been proposed in the context of domain adaptation, such as the (Ben-David et al., 2007) and discrepancy distancesDespite its name, this discrepancy is not a distance in general. (Mansour et al., 2009). These discrepancies depend on a loss function and hypothesis (i. e., predictor) class, and quantify dissimilarity through a supremum over this function class. The latter discrepancy in particular has proven remarkably useful for proving generalization bounds for adaptation (Cortes & Mohri, 2011), and while it can be estimated from samples, bounding the approximation quality relies on quantities like the VC-dimension of the hypothesis class, which might not be always known or easy to compute.
Dataset Distance via Parameter Sensitivity
The Fisher information metric is a classic notion from information geometry (Amari, 1985; Amari & Nagaoka, 2000) that characterizes a parametrized probability distribution locally through the sensitivity of its density to changes in the parameters. In machine learning, it has been used to analyze and improve optimization approaches (Amari, 1998) and to measure the capacity of neural networks (Liang et al., 2019). In recent work, Achille et al. (2019) use this notion to construct vector representations of tasks, which they then use to define a notion of similarity between these. They show that this notion recovers taxonomic similarities and is useful in meta-learning to predict whether a certain feature extractor will perform well in a new task. While this notion shares with ours its agnosticism of the number of classes and their semantics, it differs in the fact that it relies on a probe network trained on a specific dataset, so its geometry is heavily influenced by the characteristics of this network. Besides the Fisher information, a related information-theoretic notion of complexity that can be used to characterize tasks is the Kolmogorov Structure Function (Li, 2006), which Achille et al. (2018) use to define a notion of reachability between tasks.
Optimal Transport-based distributional distances
The general idea of representing complex objects via distributions, which are then compared through optimal transport distances, is an active area of research. Also driven by the appeal of their closed-form Wasserstein distance, Muzellec & Cuturi (2018) propose to embed objects as elliptical distributions, which requires differentiating through these distances, and discuss various approximations to scale up these computations. Frogner et al. (2019) extend this idea but represent the embeddings as discrete measures (i. e., point clouds) rather than Gaussian/Elliptical distributions. Both of these works focus on embedding and consider only within-dataset comparisons. Also within this line of work, Delon & Desolneux (2019) introduce a Wasserstein-type distance between Gaussian mixture models. Their approach restricts the admissible transportation couplings themselves to be Gaussian mixture models, and does not directly model label-to-label similarity. More generally, the Gromov-Wasserstein distance (Mémoli, 2011) has been proposed to compare collections across different domains (Mémoli, 2017; Alvarez-Melis & Jaakkola, 2018), albeit leveraging only features, not labels.
Hierarchical OT distances
The distance we propose can be understood as a hierarchical OT distance, i. e., one where the ground metric itself is defined through an OT problem. This principle has been explored in other contexts before. For example, Yurochkin et al. (2019) use a hierarchical OT distance for document similarity, defining a inner-level distance between topics and a outer-level distance between documents using OT. (Dukler et al., 2019) on the other hand use a nested Wasserstein distance as a loss for generative model training, motivated by the observation that the Wasserstein distance is better suited to comparing images than the usual pixel-wise metric used as ground metric. Both the goal, and the actual metric, used by these approaches differs from ours.
Optimal Transport for Domain Adaptation
Using label information to guide the optimal transport problem towards class-coherent matches has been explored before, e. g., by enforcing group-norm penalties (Courty et al., 2017) or through submodular cost functions (Alvarez-Melis et al., 2018). These works are focused on the unsupervised domain adaptation setting, so their proposed modifications to the OT objective use only label information from one of the two domains, and even then, do so without explicitly defining a metric between these. Furthermore, they do not lead to proper distances, and these works deal with a single static pair of tasks, so they lack analysis of the distance across multiple source and target datasets.
Background on Optimal Transport
Optimal transport (OT) is a powerful and principled approach to compare probability distributions, with deep roots in statistics, computer science and applied mathematics (Villani, 2003; 2008). Among many desirable properties, these distances leverage the geometry of the underlying space, making them ideal for comparing distributions, shapes and point clouds (Peyré & Cuturi, 2019).
The OT problem considers a complete and separable metric space , along with probability measures and . These can be continuous or discrete measures, the latter often used in practice as empirical approximations of the former whenever working in the finite-sample regime. The Kantorovich formulation Kantorovitch (1942) of the transportation problem reads:
Whenever is equipped with a metric , it is natural to use it as ground cost, e. g., for some . In such case, is called the -Wasserstein distance. The case is also known as the Earth Mover’s Distance (Rubner et al., 2000).
The measures and are rarely known in practice. Instead, one has access to finite samples . In that case, one can construct discrete measures and , where , are vectors in the probability simplex, and the pairwise costs can be compactly represented as an matrix , i. e., . In this case, Equation (1) becomes a linear program. Solving this problem scales cubically on the sample sizes, which is often prohibitive in practice. Adding an entropy regularization, namely
where is the relative entropy, leads to a problem that can be solved much more efficiently (Cuturi, 2013; Altschuler et al., 2017) and with better sample complexity (Genevay et al., 2019) than the original one. The Sinkhorn divergence (Genevay et al., 2018), defined as
has various desirable properties, e. g., it is positive, convex and metrizes the weak∗ convergence of distributions (Feydy et al., 2019).
Optimal Transport between Datasets
The definition of dataset is notoriously inconsistent across the machine learning literature, sometimes referring only to features or both features and labels. In the context of supervised learning, where the ultimate goal is to estimate predictors (or conditional distributions ), we define a dataset as a set of feature-label pairs over a certain feature space and label set . For simplicity, we will use to denote these pairs, and for their underlying space.
Henceforth, we focus on the case of classification, so shall be a finite set. We consider two datasets and , and assume, for simplicity, that their feature spaces have the same dimensionality, but will discuss how to relax this assumption later on. On the other hand, we make no assumptions on the label sets and whatsoever. In particular, the classes these encode could be partially overlapping or related (e. g., imagenet and cifar-10) or completely disjoint (e. g., cifar-10 and mnist). Although not a formal assumption of our approach, it will be useful to think of the samples in these two datasets as being drawn from joint distributions and .
Given and , our goal is to define a distance that depends exclusively on the information contained in these datasets. The probabilistic interpretation of these collections suggests a simple-yet-proven approach: comparing these datasets by means of a statistical divergence on their joint distributions. Among many such notions, optimal transport stands out because of various characteristics described in Section 3: its direct use of the geometry of the underlying space, its characterization of distance as correspondence (which will prove to have various useful applications in this context) and the vast theory, spanning three centuries, which it is built upon.
for is a metric on .
In most applications, is readily available, e. g., as the euclidean distance in the feature space. On the other hand, will rarely be so, particularly between labels from unrelated label sets (e. g., between cars in one image domain and and dogs in the other). If we had some prior knowledge of the label spaces, we could use it to define a notion of distance between pairs of labels. However, in the challenging —but common— case where no such knowledge is available, the only information we have about the labels is their occurrence in relation to the feature vectors . Thus, we can take advantage of the fact that we have a meaningful metric in and use it to compare labels. Arguably, the simplest such approach is as follows. Let us define , i. e., is the set of feature vectors with label in dataset , and let be its cardinality. With this, a distance between two labels and can be defined as the distance between the centroids of their associated feature vector collections:
Although appealing for its simplicity, representing the collections only through their mean is too simplistic for real datasets. Ideally, we would like to represent labels through the actual distribution over the feature space that they define, namely, by means of the map , of which can be understood as a finite sample. If we use this representation, defining a distance between labels boils down to choosing a statistical divergence between their associated distributions. Once more, there are many possible choices for this distance, but —yet again— we argue that an OT is an ideal choice, since the notion of divergence we seek should: (i) provide a valid metric, (ii) be computable from finite samples, which is crucial since the distributions are not available in analytic form, and (iii) be able to deal with sparsely-supported distributions, all of which OT satisfies.
The approach described so far grounds the comparison of the distributions to the feature space , so we can simply use as the optimal transport cost, leading to a p-Wasserstein distance between labels: , and in turn, to the following distance between feature-label pairs:
This gives us a point-wise notion of distance in , but we ultimately seek a distance between distributions over this space, i. e., between joint distributions . Optimal transport allows us to lift the ground (i. e., point-wise) metric defined above into a distance between measures:
The following result, an immediate consequence of the discussion above, states that Eq. (6) is a proper distance – the Optimal Transport Dataset Distance (otdd).
defines a valid metric on the space of measures over feature and label-distribution pairs.
It remains to describe how the distributions are to be represented. A flexible non-parametric approach would be to treat the samples in as support points of a uniform empirical measure, i. e., , as described in Section 3. The main downside of this approach is that each evaluation of (5) involves solving an optimization problem, which could be prohibitive. Indeed, in Section C.1 we show that for datasets of size , this approach has worst-case complexity.
With this, we model each label-feature distribution as a Gaussian Distribution whose parameters are the sample mean and covariance of .
The main motivation behind this choice is that the 2-Wasserstein distance between Gaussian distributions and has as an analytic form:
where denotes the matrix square root. Furthermore, whenever and commute, this further simplifies to
When using Eq. (7) in the point-wise distance (5), we denote the resulting distance (6) by .
Representing label-defined distributions as Gaussians might seem like a heuristic choice driven only by algebraic convenience. However, the following result, a consequence of a bound by Gelbrich (1990), shows that this approximation lower-bounds the distance that would be obtained had it been computed using the label distances on the true distributions (regardless of their form):
For any two datasets , we have:
Furthermore, if the label distributions are all Gaussian or elliptical, these quantities are equal, i. e., is exact.
An illustration of the OTDD in a synthetic dataset summarizing its main characteristics is shown in Figure 2.
Computational Considerations
Since our goal in this work is to use the proposed dataset distance as a tool for tasks like transfer learning in realistic (i. e., large) machine learning datasets, scalability is crucial. Indeed, most compelling use cases of any notion of distance between datasets will involve computing it repeatedly on very large samples.
While estimation of Wasserstein —and more generally, optimal transport— distances is known to be computationally expensive in general, in Section 3 we briefly discussed how entropy regularization can be used to trade-off accuracy for runtime. Recall that both the general and Gaussian versions of the dataset distance proposed in Section 4 involve solving optimal transport problems (though the latter, owing the closed form solution of subproblem (7), only requires optimization for the global problem). Therefore, both of these distances benefit from approximate OT solvers.
The steps we propose next are motivated by the observation that, unlike traditional OT distances for which the cost of computing pair-wise distance is negligible compared to the complexity of the optimization routine, in our case the latter dominates, since it involves computing multiple OT distances itself. In order to speed up computation, we first precompute and store in memory all label-to-label pairwise distances , and retrieve them on-demand during the optimization of the global OT problem.
For , computing the label-to-label distances is dominated by the cost of computing matrix square roots, which if done exactly involves a full eigendecomposition. Instead, it can be computed approximately using the Newton-Schulz iterative method (Higham, 2008; Muzellec & Cuturi, 2018). Besides runtime, loading all examples of a given class to memory (to compute means and covariances) might be infeasible for large datasets (especially if running on GPU), so we instead use a two-pass stable online batch algorithm to compute these statistics (Chan et al., 1983).
The following result summarizes the time complexity of our two distances and sheds light on the trade-off between precision and efficiency they provide.
For datasets of size and , with and classes, dimension , and maximum class size , both and incur in a cost of for solving the global OT problem -approximately, while the worst-case complexity for computing the label-to-label pairwise distances (5) is O\bigl{(}nm(d+\mathfrak{n}^{3}\log\mathfrak{n}+d\mathfrak{n}^{2})\bigr{)} for and O\bigl{(}nmd+pqd^{3}+d^{2}\mathfrak{n}(p+q)\bigr{)} for .
In most practical applications, the cost of computing pairwise distances will dominate, making superior. For example, if and the largest class size is , this step becomes —prohibitive for all but toy datasets— for but only for .
Experiments
A driving motivation for proposing a dataset distance was to provide a learning-free criterion on which to select a source dataset for transfer learning. In this section, we put this hypothesis to test on a simple domain adaptation setting on mnist (LeCun et al., 2010) and three of its extensions: fashion-mnist (Xiao et al., 2017), kmnist (Clanuwat et al., 2018) and the letters split of emnist (Cohen et al., 2017), in addition to usps. All datasets consist of 10 classes, except emnist, for which the selected split has 26 classes. Throughout this section, we use a simple LeNet-5 neural network (two convolutional layers, three fully conntected ones) with ReLU activations. When carrying out adaptation, we freeze the convolutional layers and fine-tune only the top three layers.
We first compute all pairwise OTDD distances (Fig 4). For the example of , Figure 3 illustrates two key components of the computation of the distance: the label-to-label distances (left) and the optimal coupling obtained for two choices of entropy regularization parameter (center, right). The diagonal elements of the first plot (i. e., distances between corresponding digit classes) are overall relatively smaller than off-diagonal elements. Interestingly, the 0 class of usps appears remarkably far from all mnist digits under this metric. On the other hand, most correspondences lie along the (block) diagonal of , which shows the dataset distance is able to infer class-coherent correspondences across them.
We test the robustness of the distance by computing it repeatedly for varying sample sizes. The results (Fig. 9, Appendix F) show that the distance converges towards a fixed value as sample sizes grow, but interestingly, small sample sizes for usps lead to wider variability, suggesting that this dataset itself is more heterogeneous than mnist.
Despite both consisting of digits, mnist and usps are not the closest among these datasets according to the OTDD, as Figure 4 shows. The closest pair is instead (mnist, emnist), while fashion-mnist appears comparatively far from all others, particularly mnist.
Next, we compare these distances against the transferability between datasets, i. e., the gain in performance from using a model pertrained on the source domain and fine-tuning it on the target domain. To make these numbers comparable across adaptation pairs which involve datasets of very different hardness, we define the transferability of a source domain to a target domain as the relative decrease in classification error when doing adaptation compared to training only on the target domain, i. e.,
We run the adaptation task 10 times with different random seeds for each pair of datasets, and compare against their distance. The strong significant correlation between these (Fig. 5) shows that the OTDD is highly predictive of transferability across these datasets. In particular, emnist led to the best adaptation to mnist, justifying the —initially counter-intuitive— value of the OTDD.
2 Distance-Driven Data Augmentation
Data augmentation —i. e., applying carefully chosen transformations on a dataset to enhance its quality and diversity— is another key aspect of transfer learning that has substantial empirical effect on the quality of the transferred model yet lacks principled guidelines. Here, we investigate if the OTDD could be used to compare and select among possible augmentations.
For a fixed source-target dataset pair, we generate replicas of the source data with various transformations applied to it, compute their distance to the target dataset, and compare against the transferability as before. We present results for a small-scale (mnist usps) and a larger-scale (Tiny-ImageNetcifar-10) setting. The transformations we use on mnist consist of rotations by a fixed degree , random rotations , random affine transformations, center-crops and random crops. For Tiny-ImageNet we randomly vary brightness, contrast, hue and saturation. The models use are respectively the LeNet-5 and a ResNet-50 (training details provided in Appendix E).
The results in both of these settings (Figures 6 and 7) show, again, a strong significant correlation between these two. A reader familiar with the mnist and usps datasets will not be surprised by the fact that cropping images from the former leads to substantially better performance on the latter, while most rotations degrade transferability.
3 Transfer Learning for Text Classification
Natural Language Processing (NLP) is of the areas where large-scale transfer learning has had the most profound impact over the past few years, in part driven by the availability of off-the-shelf large language-models pretrained on massive amounts of the data (Peters et al., 2018; Devlin et al., 2019; Radford et al., 2019).
While natural language inherently lacks the fixed-size continuous vector representation required by our framework to compute pointwise distances, we can take advantage of precisely these pretrained models to embed sentences in vector space, furnishing them with a rich geometry. In our experiments, we first embed every sentence of every dataset using the (base) bert model (Devlin et al., 2019),Using the sentence_transfomers library. and then compute OTDD on these embedded datasets.
As before, we simulate a challenging adaptation setting by keeping only 100 examples per target class. For every pair of datasets, we first fine-tune the bert model using the entirety of the source domain data, after which we fine-tune and evaluate on the target domain. Figure 8 shows that the OT dataset distance is highly correlated with transferability in this setting too. Interestingly, adaptation often leads to drastic degradation of performance in this case, which suggests that off-the-shelf bert is on its own powerful and flexible enough to initialize many of these tasks, and therefore choosing the wrong domain for initial training might destroy some of that information.
Discussion
We have shown that the notion of distance between datasets proposed in this work is scalable and flexible enough to be used in realistic transfer learning scenarios, all the while offering appealing theoretical properties, interpretable comparisons and requiring minimal assumptions on the underlying datasets.
There are many natural extensions of this framework. Here we assumed that the datasets where defined on feature spaces of the same dimension, but one could instead leverage a relational notion such as the Gromov-Wasserstein distance (Mémoli, 2011) to compute the distance between datasets whose features and not directly comparable. On the other hand, our efficient implementation relies on modeling groups of points with the same label as Gaussian distributions. This could naturally be extended to more general distributions for which the Wasserstein distance either has an analytic solution or at least can be computed efficiently, such as elliptic distributions (Muzellec & Cuturi, 2018), Gaussian mixture models (Delon & Desolneux, 2019), certain Gaussian Processes (Mallasto & Feragen, 2017), or tree metrics (Le et al., 2019).
In this work, we purposely excluded two key aspects of any learning task from our notion of distance: the loss function and the predictor function class. While we posit that it is crucial to have a notion of distance that is independent of these choices, it is nevertheless appealing to ask whether our distance could be extended to take those into account, ideally involving minimal training. Exploring different avenues to inject such information into this framework will be the focus of our future work.
References
Appendix A Proof of Proposition 4.1
Whenever the cost function used in the of optimal transport problem is a metric in a given space , the optimal transport problem is a distance (the Wasserstein distance) on (Villani, 2008, Chapter 6). Therefore, it suffices to show that the cost function defined in Eq. (5) is indeed a distance. Clearly, it is symmetric because both and are. In addition, since both of these are distances:
where the last step is an application of Minkowski’s inequality. Hence, satisfies the triangle inequality, and therefore it is a metric on . We therefore conclude that the value of the optimal transport (6) that uses this metric as a cost function is a distance itself. ∎
Appendix B Proof of Proposition 4.2
Our proof relies directly on a well-known bound for the 2-Wasserstein distance between distributions by (Gelbrich, 1990):
where \textup{W}_{2}^{2}\bigl{(}\mathcal{N}(\mu_{\alpha},\Sigma_{\alpha}),\mathcal{N}(\mu_{\beta},\Sigma_{\beta})) is as in Eq. (7).
In the notation of Section 3, Lemma B.1 implies that for every feature-label pairs and , we have:
for every coupling . In particular, for the minimizing , we obtain that
Clearly, Gelbrich’s bound holds with equality when and are indeed Gaussian. More generally, equality is attained for elliptical distributions with the same density generator (Kuhn et al., 2019)). This immediately implies equality of the two quantities in equation (13) in that case. ∎
Appendix C Time Complexity Analysis
Direct computation of the distance (5) involves two main steps:
computing pairwise pointwise distances (each requiring solution of a label-to-label OT sub-problem), and
a global OT problem between the two samples.
Step (ii) is identical for both the general distance and its Gaussian approximation counterpart , so we analyze it first. This is an OT problem between two discrete distributions of size and , which can be solved exactly in O\bigl{(}(n+m)nm\log(nm)\bigr{)} using interior point methods or Orlin’s algorithm for the uncapacitated min cost flow problem (Peyré & Cuturi, 2019). Alternatively, it can be solved -approximately in time using the Sinkhorn algorithm (Altschuler et al., 2017).
We next analyze step (i) individually for the two OTDD versions. Combined, they provide a proof of Theorem 5.1.
Consider a single pair of points, and . Evaluating has complexity, while is an OT problem which itself requires computing a distance matrix (at cost ), and then solving the OT problem, which as discussed before, be done exactly in O\bigl{(}(n_{s}^{i}+n_{t}^{j})n_{s}^{i}n_{t}^{j}\log(n_{s}^{i}+n_{t}^{j})\bigr{)} or -approximately in .
For simplicity, let us denote , and the size of the largest label cluster in each dataset, and the overall largest one. Using these, and combining all of the above, the overall worst case complexity for the computation of the pairwise distances can be expressed as
As before, consider a pair of points and whose cluster sizes are and respectively. As mentioned in Section 5, for we first compute all the per-class means and covariance matrices. This step is clearly dominated by latter, which is .technically, this would be where is the coefficient of matrix multiplication, but we take for simplicity. Considering all labels from both datasets, this amounts to a worst-case complexity of O\bigl{(}d^{2}(k_{s}\mathfrak{n}_{s}+k_{t}\mathfrak{n}_{t})\bigr{)}.
Once the means and covariances have been computed, we precompute all the pair-wise label-to-label distances using Eq. (7). This computation is dominated by the matrix square roots. If done exactly, these involve a full eigendecomposition, at cost , so the total cost for this step is .
Finally, while computing the pairwise distance, we will incur in to obtain . Putting all of these together, and replacing by , we obtain a total cost for precomputing all the point-wise distances of:
Appendix D Dataset Details
Information about all the datasets used, including references, are provided in Table LABEL:tab:dataset_details.
Appendix E Optimization and Training Details
For the Tiny-ImageNet to Cifar-10 adaptation results, we use a ResNet-50 trained for 300 epochs using SGD with learning rate 0.1 momentum 0.9 and weight decay It was fine-tuned for 30 epochs on the target domain using SGD with same parameters except 0.01 learning rate. We discard pairs for which the variance on adaptation accuracy is beyond a certain threshold.
Our implementation of the OTDD relies on the potpot.readthedocs.io/en/stable/ and geomlosswww.kernel-operations.io/geomloss/ python packages.