Learning Optimal Transport Plans Via Autoregressive Token Regression
Abstract
Optimal Transport (OT) provides a principled framework for comparing probability distributions. Classical solvers remain the gold standard for this task, providing exact solutions. At the same time, OT is an ideal testbed for studying whether learning-based amortized models can internalize the structure of OT type problems. Here, ground-truth solutions are available, constraints are explicit, and exact discrete solutions have well-characterized sparsity. We describe ToR-OT, a transformer-based model that learns to predict discrete OT plans through autoregressive token regression. We reframe OT into a structured sequence generation task. ToR-OT decodes sparse transport plans as sequences rather than solving an optimization task at test time. We focus on three properties: (a) Zero-shot generalization: pretraining on a large synthetic corpus of OT instances allows a single model to generalize to unseen distribution pairs and heterogeneous cost metrics without retraining. (b) Sparse primal decoding: by exploiting the inherent sparsity in exact discrete OT solutions, the decoder predicts only nonzero entries (avoids dense coupling outputs). (c) Amortized inference: on low-dimensional discrete problems, ToR-OT gives accurate approximate plans and compares favorably with per-instance neural OT baselines. Overall, ToR-OT is not yet a replacement for classical solvers, but best viewed as a step toward understanding how autoregressive transformers can learn structure of constrained optimization problems like OT.