Multi-Marginal Inverse Optimal Transport for Contrastive Learning Via Explicit Anchor-Positive-Negative Coupling
Organizations: Department of Electrical and Computer Engineering Tufts University · Department of Engineering, Engineering Technology East Tennessee State University · Department of Electrical and Computer Engineering Boston University
Abstract
Inverse Optimal Transport (OT) based methods for representation learning learn representations such that the global OT coupling between a pair of data marginals in the representation space, concentrates on the positive pairs. This is in contrast to previous methods that primarily focused on pairwise matching. However, these methods utilize negative pairs and hence are not truly contrastive in their approach. We show that this leads to issues of dimensional collapse and hence degraded downstream performance. To alleviate this, we develop a novel multi-marginal (MM) inverse OT (IOT) contrastive learning (CL) approach called Neg-MMIOT-CL, which learns representations such that the global multi-marginal OT (MMOT) coupling between a triple of data marginals, with respect to a carefully designed ground-cost between triplets of data points in the representation space, concentrates on the anchor-positive-negative . For a latent class model, we empirically show that Neg-MMIOT-CL alleviates dimensional collapse. Furthermore, for a specific choice of ground cost for all triplets in representation space, we prove that the optimal representation configuration for Neg-MMIOT-CL exhibits equiangular property for within-class and across-class representations, which translates to Neural-Collapse when the representation dimension is larger than the number of classes minus one -- a result that is for pairwise contrastive learning methods. Finally, we propose Neg-IOT-CL-PushPull, that is a computationally efficient alternative to Neg-MMIOT-CL, alleviating the high cost of computing MMOT plans needed during implementation. We apply these methods on both synthetic and real-world datasets and show significant improvements over existing OT-based contrastive learning methods.
Figures & tables
| Setting | Method | MNIST | SVHN | CIFAR-10 | CIFAR-100 | Tiny-ImageNet | |||||
| ResNet50 | ViT-B/16 | ResNet50 | ViT-B/16 | ResNet50 | ViT-B/16 | ResNet50 | ViT-B/16 | ResNet50 | ViT-B/16 | ||
| SCL | InfoNCE | 99.36 | 99.49 | 91.96 | 94.07 | 91.39 | 90.40 | 74.01 | 73.59 | 62.60 | 46.32 |
| InvaSpread | 98.98 | 99.61 | 81.79 | 85.35 | 78.46 | 72.62 | 63.26 | 67.08 | 55.06 | 44.19 | |
| Standard OT | 99.21 | 99.60 | 92.08 | 94.24 | 90.78 | 84.72 | 72.71 | 73.49 | 61.20 | 45.27 | |
| Neg-IOT-CL-PushPull (Ours) | 99.30 | 99.64 | 91.89 | 94.38 | 89.59 | 88.69 | 74.80 | 74.67 | 62.04 | 46.54 | |
| Neg-MMIOT-CL (Ours) | 99.41 | 99.66 | 92.74 | 94.87 | 90.64 | 91.97 | 75.17 | 74.62 | 63.54 | 48.02 | |
| CLIP-Loss | Image Text | Text Image | CIFAR-10 | CIFAR-100 | ||||
|---|---|---|---|---|---|---|---|---|
| Top-1 | Top-5 | Top-1 | Top-5 | Top-1 | Top-5 | Top-1 | Top-5 | |
| infoNCE | 67.22 | 92.51 | 65.91 | 94.61 | 23.37 | 74.82 | 6.36 | 27.59 |
| Standard OT | 43.92 | 81.86 | 30.90 | 70.54 | 26.62 | 78.56 | 4.77 | 22.95 |
| Standard OT - Uniform | 64.25 | 92.40 | 63.22 | 92.06 | 25.97 | 78.67 | 6.51 | 27.68 |
| DBOT | 15.89 | 46.68 | 15.52 | 45.92 | 24.98 | 70.36 | 5.72 | 27.19 |
| Fused-Gromov | 14.32 | 43.92 | 13.94 | 42.87 | 24.99 | 76.06 | 3.72 | 19.40 |
Appendix figures & tables21 assets
Supplementary material from the paper’s appendix.
Appendix
| SCL | MNIST | |||||||
|---|---|---|---|---|---|---|---|---|
| ResNet18 | ResNet34 | ResNet50 | ViT-B/16 | |||||
| Linear | k-NN | Linear | k-NN | Linear | k-NN | Linear | k-NN | |
| infoNCE | ||||||||
| InvaSpread | ||||||||
| Standard OT | ||||||||
| Neg-IOT-CL-PushPull (Ours) | ||||||||
| SCL | SVHN | |||||||
|---|---|---|---|---|---|---|---|---|
| ResNet18 | ResNet34 | ResNet50 | ViT-B/16 | |||||
| Linear | k-NN | Linear | k-NN | Linear | k-NN | Linear | k-NN | |
| infoNCE | ||||||||
| InvaSpread | ||||||||
| Standard OT | ||||||||
| Neg-IOT-CL-PushPull (Ours) | ||||||||
| SCL | CIFAR-10 | |||||||
|---|---|---|---|---|---|---|---|---|
| ResNet18 | ResNet34 | ResNet50 | ViT-B/16 | |||||
| Linear | k-NN | Linear | k-NN | Linear | k-NN | Linear | k-NN | |
| infoNCE | ||||||||
| InvaSpread | ||||||||
| Standard OT | ||||||||
| Neg-IOT-CL-PushPull (Ours) | ||||||||
| SCL | CIFAR-100 | |||||||
|---|---|---|---|---|---|---|---|---|
| ResNet18 | ResNet34 | ResNet50 | ViT-B/16 | |||||
| Linear | k-NN | Linear | k-NN | Linear | k-NN | Linear | k-NN | |
| infoNCE | ||||||||
| InvaSpread | ||||||||
| Standard OT | ||||||||
| Neg-IOT-CL-PushPull (Ours) | ||||||||
| SCL | Tiny-ImageNet | |||||||
|---|---|---|---|---|---|---|---|---|
| ResNet18 | ResNet34 | ResNet50 | ViT-B/16 | |||||
| Linear | k-NN | Linear | k-NN | Linear | k-NN | Linear | k-NN | |
| infoNCE | ||||||||
| InvaSpread | ||||||||
| Standard OT | ||||||||
| Neg-IOT-CL-PushPull (Ours) | ||||||||
| UCL | MNIST | |||||||
|---|---|---|---|---|---|---|---|---|
| ResNet18 | ResNet34 | ResNet50 | ViT-B/16 | |||||
| Linear | k-NN | Linear | k-NN | Linear | k-NN | Linear | k-NN | |
| infoNCE | ||||||||
| InvaSpread | ||||||||
| Standard OT | ||||||||
| Neg-IOT-CL-PushPull (Ours) | ||||||||
| UCL | SVHN | |||||||
|---|---|---|---|---|---|---|---|---|
| ResNet18 | ResNet34 | ResNet50 | ViT-B/16 | |||||
| Linear | k-NN | Linear | k-NN | Linear | k-NN | Linear | k-NN | |
| infoNCE | ||||||||
| InvaSpread | ||||||||
| Standard OT | ||||||||
| Neg-IOT-CL-PushPull (Ours) | ||||||||
| UCL | CIFAR-10 | |||||||
|---|---|---|---|---|---|---|---|---|
| ResNet18 | ResNet34 | ResNet50 | ViT-B/16 | |||||
| Linear | k-NN | Linear | k-NN | Linear | k-NN | Linear | k-NN | |
| infoNCE | ||||||||
| InvaSpread | ||||||||
| Standard OT | ||||||||
| Neg-IOT-CL-PushPull (Ours) | ||||||||
| UCL | CIFAR-100 | |||||||
|---|---|---|---|---|---|---|---|---|
| ResNet18 | ResNet34 | ResNet50 | ViT-B/16 | |||||
| Linear | k-NN | Linear | k-NN | Linear | k-NN | Linear | k-NN | |
| infoNCE | ||||||||
| InvaSpread | ||||||||
| Standard OT | ||||||||
| Neg-IOT-CL-PushPull (Ours) | ||||||||
| UCL | Tiny-ImageNet | |||||||
|---|---|---|---|---|---|---|---|---|
| ResNet18 | ResNet34 | ResNet50 | ViT-B/16 | |||||
| Linear | k-NN | Linear | k-NN | Linear | k-NN | Linear | k-NN | |
| infoNCE | ||||||||
| InvaSpread | ||||||||
| Standard OT | ||||||||
| Neg-IOT-CL-PushPull (Ours) | ||||||||
| Learning Rate | Linear | k-NN |
|---|---|---|
| 79.52 | 79.40 | |
| 82.00 | 81.92 | |
| 82.15 | 82.66 | |
| 0.01 | 80.00 | 79.85 |
| 0.03 | 77.70 | 77.81 |
| 0.1 | 75.95 | 75.04 |
| CIFAR-100 | Tiny-ImageNet | |||||||
|---|---|---|---|---|---|---|---|---|
| SCL | ResNet18 | ResNet34 | ResNet50 | ViT-B/16 | ResNet18 | ResNet34 | ResNet50 | ViT-B/16 |
| InfoNCE | 3.22 | 4.40 | 5.66 | 13.23 | 6.44 | 9.37 | 11.08 | 27.68 |
| Standard OT | 3.40 | 4.86 | 6.46 | 13.59 | 6.92 | 9.35 | 11.56 | 28.39 |
| InvaSpread | 3.16 | 4.48 | 5.85 | 13.02 | 6.87 | 9.49 | 11.81 | 28.20 |
| OT-PushPull | 3.69 | 5.62 | 7.04 | 13.87 | 7.69 | 10.19 | 12.73 | 29.64 |
| MMIOT | 6.01 | 7.04 | 9.64 | 16.71 | 10.25 | 13.39 | 16.05 | 35.19 |
| CIFAR-100 | Tiny-ImageNet | |||||||
|---|---|---|---|---|---|---|---|---|
| UCL | ResNet18 | ResNet34 | ResNet50 | ViT-B/16 | ResNet18 | ResNet34 | ResNet50 | ViT-B/16 |
| InfoNCE | 71.76 | 74.05 | 75.55 | 86.44 | 244.79 | 250.85 | 277.16 | 288.26 |
| Standard OT | 72.38 | 76.26 | 76.37 | 86.80 | 231.37 | 282.39 | 351.46 | 286.67 |
| InvaSpread | 75.21 | 80.41 | 83.78 | 87.61 | 257.83 | 324.82 | 360.12 | 287.85 |
| OT-PushPull | 73.54 | 76.28 | 77.34 | 87.89 | 294.22 | 283.76 | 351.88 | 288.30 |
| MMIOT | 79.28 | 82.51 | 82.30 | 100.57 | 273.87 | 294.03 | 376.97 | 300.93 |
Explore similar work
Positive Pair Geometry Matters: Optimal Transport for Contrastive Learning of Visual Representations
Optimal VC Dimension of Contrastive Learning with Margin
anchor--positive--negative'' triplets $(i,j^{+},k^{-})$, indicating that item is closer to than to .'' Despite its success, understanding why contrastive learning leads to representations of high \textit{generalization} quality---beyond the often pessimistic predictions from PAC-learning---remains a central question. Recently, \citet*{alon2024optimal} proved that, for PAC-learning -dimensional Euclidean representations of -point datasets, triplets are necessary and sufficient, while they posed as an open question whether their VC dimension bounds for the more realistic setting of \textit{contrastive learning with a margin} can be improved. For a margin parameter , a triplet is satisfied by the embedding , if . In this work, we resolve their question by proving that the VC dimension of contrastive learning under any margin is in fact , improving on the previous bound of . We also establish that the bounds are optimal up to constant factors, by providing a matching lower bound of (the previously known lower bound was ), for .