flash-sinkhorn 0.3.2__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (48) hide show
  1. flash_sinkhorn-0.3.2/LICENSE +21 -0
  2. flash_sinkhorn-0.3.2/PKG-INFO +315 -0
  3. flash_sinkhorn-0.3.2/README.md +283 -0
  4. flash_sinkhorn-0.3.2/pyproject.toml +57 -0
  5. flash_sinkhorn-0.3.2/setup.cfg +4 -0
  6. flash_sinkhorn-0.3.2/src/flash_sinkhorn/__init__.py +38 -0
  7. flash_sinkhorn-0.3.2/src/flash_sinkhorn/_autograd.py +519 -0
  8. flash_sinkhorn-0.3.2/src/flash_sinkhorn/bench/__init__.py +3 -0
  9. flash_sinkhorn-0.3.2/src/flash_sinkhorn/bench/bench_backward.py +1091 -0
  10. flash_sinkhorn-0.3.2/src/flash_sinkhorn/bench/bench_forward.py +1715 -0
  11. flash_sinkhorn-0.3.2/src/flash_sinkhorn/cg.py +369 -0
  12. flash_sinkhorn-0.3.2/src/flash_sinkhorn/hvp.py +835 -0
  13. flash_sinkhorn-0.3.2/src/flash_sinkhorn/implicit_grad.py +315 -0
  14. flash_sinkhorn-0.3.2/src/flash_sinkhorn/kernels/__init__.py +92 -0
  15. flash_sinkhorn-0.3.2/src/flash_sinkhorn/kernels/_common.py +231 -0
  16. flash_sinkhorn-0.3.2/src/flash_sinkhorn/kernels/_triton_helpers.py +174 -0
  17. flash_sinkhorn-0.3.2/src/flash_sinkhorn/kernels/apply_flash.py +951 -0
  18. flash_sinkhorn-0.3.2/src/flash_sinkhorn/kernels/apply_ott.py +566 -0
  19. flash_sinkhorn-0.3.2/src/flash_sinkhorn/kernels/sinkhorn_flashstyle_sqeuclid.py +1373 -0
  20. flash_sinkhorn-0.3.2/src/flash_sinkhorn/kernels/sinkhorn_triton_apply_fused_sqeuclid.py +486 -0
  21. flash_sinkhorn-0.3.2/src/flash_sinkhorn/kernels/sinkhorn_triton_apply_sqeuclid.py +17 -0
  22. flash_sinkhorn-0.3.2/src/flash_sinkhorn/kernels/sinkhorn_triton_cg_dense.py +582 -0
  23. flash_sinkhorn-0.3.2/src/flash_sinkhorn/kernels/sinkhorn_triton_cg_python_batched.py +399 -0
  24. flash_sinkhorn-0.3.2/src/flash_sinkhorn/kernels/sinkhorn_triton_geomloss_sqeuclid.py +310 -0
  25. flash_sinkhorn-0.3.2/src/flash_sinkhorn/kernels/sinkhorn_triton_grad_sqeuclid.py +1207 -0
  26. flash_sinkhorn-0.3.2/src/flash_sinkhorn/kernels/sinkhorn_triton_ott_sqeuclid.py +750 -0
  27. flash_sinkhorn-0.3.2/src/flash_sinkhorn/samples_loss.py +597 -0
  28. flash_sinkhorn-0.3.2/src/flash_sinkhorn/sinkhorn_solvers.py +707 -0
  29. flash_sinkhorn-0.3.2/src/flash_sinkhorn/testing/__init__.py +3 -0
  30. flash_sinkhorn-0.3.2/src/flash_sinkhorn/testing/reference_hvp.py +356 -0
  31. flash_sinkhorn-0.3.2/src/flash_sinkhorn/testing/reference_sinkhorn.py +233 -0
  32. flash_sinkhorn-0.3.2/src/flash_sinkhorn/testing/test_apply_plan_flashstyle.py +585 -0
  33. flash_sinkhorn-0.3.2/src/flash_sinkhorn/testing/test_flashstyle_parity.py +679 -0
  34. flash_sinkhorn-0.3.2/src/flash_sinkhorn/testing/test_geomloss_sinkhorn_triton.py +411 -0
  35. flash_sinkhorn-0.3.2/src/flash_sinkhorn/testing/test_geomloss_vs_triton.py +68 -0
  36. flash_sinkhorn-0.3.2/src/flash_sinkhorn/testing/test_half_cost.py +143 -0
  37. flash_sinkhorn-0.3.2/src/flash_sinkhorn/testing/test_hvp_parity.py +386 -0
  38. flash_sinkhorn-0.3.2/src/flash_sinkhorn/testing/test_hvp_sqeuclid.py +332 -0
  39. flash_sinkhorn-0.3.2/src/flash_sinkhorn/testing/test_ott_vs_triton.py +235 -0
  40. flash_sinkhorn-0.3.2/src/flash_sinkhorn/testing/test_samples_loss_api.py +355 -0
  41. flash_sinkhorn-0.3.2/src/flash_sinkhorn/testing/test_semi_unbalanced_forward.py +145 -0
  42. flash_sinkhorn-0.3.2/src/flash_sinkhorn/testing/test_sinkhorn_triton.py +152 -0
  43. flash_sinkhorn-0.3.2/src/flash_sinkhorn/testing/test_unbalanced_sinkhorn.py +564 -0
  44. flash_sinkhorn-0.3.2/src/flash_sinkhorn.egg-info/PKG-INFO +315 -0
  45. flash_sinkhorn-0.3.2/src/flash_sinkhorn.egg-info/SOURCES.txt +46 -0
  46. flash_sinkhorn-0.3.2/src/flash_sinkhorn.egg-info/dependency_links.txt +1 -0
  47. flash_sinkhorn-0.3.2/src/flash_sinkhorn.egg-info/requires.txt +11 -0
  48. flash_sinkhorn-0.3.2/src/flash_sinkhorn.egg-info/top_level.txt +1 -0
@@ -0,0 +1,21 @@
1
+ MIT License
2
+
3
+ Copyright (c) 2025 OT Triton Contributors
4
+
5
+ Permission is hereby granted, free of charge, to any person obtaining a copy
6
+ of this software and associated documentation files (the "Software"), to deal
7
+ in the Software without restriction, including without limitation the rights
8
+ to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
9
+ copies of the Software, and to permit persons to whom the Software is
10
+ furnished to do so, subject to the following conditions:
11
+
12
+ The above copyright notice and this permission notice shall be included in all
13
+ copies or substantial portions of the Software.
14
+
15
+ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
16
+ IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
17
+ FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
18
+ AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
19
+ LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
20
+ OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
21
+ SOFTWARE.
@@ -0,0 +1,315 @@
1
+ Metadata-Version: 2.4
2
+ Name: flash-sinkhorn
3
+ Version: 0.3.2
4
+ Summary: Sinkhorn optimal transport kernels in PyTorch + Triton (squared Euclidean, no cost matrix materialization).
5
+ Author: OT Triton Contributors
6
+ License-Expression: MIT
7
+ Project-URL: Homepage, https://github.com/ot-triton-lab/flash-sinkhorn
8
+ Project-URL: Repository, https://github.com/ot-triton-lab/flash-sinkhorn
9
+ Keywords: optimal-transport,sinkhorn,triton,pytorch,gpu
10
+ Classifier: Development Status :: 3 - Alpha
11
+ Classifier: Programming Language :: Python :: 3
12
+ Classifier: Programming Language :: Python :: 3.9
13
+ Classifier: Programming Language :: Python :: 3.10
14
+ Classifier: Programming Language :: Python :: 3.11
15
+ Classifier: Programming Language :: Python :: 3.12
16
+ Classifier: Topic :: Scientific/Engineering
17
+ Classifier: Environment :: GPU :: NVIDIA CUDA
18
+ Requires-Python: >=3.9
19
+ Description-Content-Type: text/markdown
20
+ License-File: LICENSE
21
+ Requires-Dist: numpy
22
+ Requires-Dist: torch>=2.5
23
+ Requires-Dist: triton>=3.1
24
+ Provides-Extra: dev
25
+ Requires-Dist: pytest; extra == "dev"
26
+ Requires-Dist: geomloss; extra == "dev"
27
+ Requires-Dist: pykeops; extra == "dev"
28
+ Requires-Dist: jax; extra == "dev"
29
+ Requires-Dist: ott-jax; extra == "dev"
30
+ Requires-Dist: matplotlib; extra == "dev"
31
+ Dynamic: license-file
32
+
33
+ <p align="center">
34
+ <img src="FlashSinkhorn.png" alt="FlashSinkhorn" width="100%">
35
+ </p>
36
+
37
+ # FlashSinkhorn
38
+
39
+ **Streaming Entropic Optimal Transport in PyTorch + Triton**
40
+
41
+ FlashSinkhorn computes Sinkhorn OT using FlashAttention-style streaming—**never materializing the n×m cost matrix**—enabling **O(nd) memory** instead of O(n²).
42
+
43
+ ## Features
44
+
45
+ - **FlashSinkhorn kernels** — shifted-potential formulation inspired by FlashAttention, 10-40% faster than previous Triton kernels at n >= 10k
46
+ - **Fused Triton kernels** for forward, gradient, and HVP
47
+ - **GeomLoss-compatible API** (`SamplesLoss`)
48
+ - **Analytic gradients** (no backprop through Sinkhorn iterations)
49
+ - **Hessian-vector products** via streaming CG solver
50
+ - **Half-cost support** (`half_cost=True`) for exact GeomLoss parity
51
+ - **Unbalanced/semi-unbalanced OT** via `reach` parameter
52
+ - **Large-D support** (d > 1024) with tiled gradient kernel
53
+ - **Early stopping** with convergence threshold
54
+
55
+ ## Install
56
+
57
+ ```bash
58
+ pip install -e .
59
+ pip install -e ".[dev]" # with dev dependencies
60
+ ```
61
+
62
+ **Requirements:** PyTorch ≥2.5, Triton ≥3.1, CUDA 12.x
63
+
64
+ ## Quick Start
65
+
66
+ ### Basic Usage
67
+
68
+ ```python
69
+ import torch
70
+ from flash_sinkhorn import SamplesLoss
71
+
72
+ x = torch.randn(4096, 64, device="cuda")
73
+ y = torch.randn(4096, 64, device="cuda")
74
+
75
+ # FlashSinkhorn is the default backend (use_flashstyle=True)
76
+ loss = SamplesLoss(loss="sinkhorn", blur=0.1, debias=True)
77
+ cost = loss(x, y)
78
+ ```
79
+
80
+ ### Gradient Flow
81
+
82
+ ```python
83
+ x = torch.randn(4096, 64, device="cuda", requires_grad=True)
84
+ y = torch.randn(4096, 64, device="cuda")
85
+
86
+ loss = SamplesLoss(loss="sinkhorn", blur=0.1, debias=True)
87
+ cost = loss(x, y)
88
+ grad_x = torch.autograd.grad(cost, x)[0] # Analytic gradient
89
+ ```
90
+
91
+ ### GeomLoss Parity
92
+
93
+ Use `half_cost=True` to match GeomLoss's cost convention:
94
+
95
+ ```python
96
+ # FlashSinkhorn with half_cost matches GeomLoss exactly
97
+ flash_loss = SamplesLoss(loss="sinkhorn", blur=0.1, half_cost=True, debias=True)
98
+
99
+ # Equivalent GeomLoss call
100
+ # geomloss_loss = geomloss.SamplesLoss(loss="sinkhorn", p=2, blur=0.1, debias="positive")
101
+ ```
102
+
103
+ ### Unbalanced OT
104
+
105
+ For distributions with different total mass or outliers:
106
+
107
+ ```python
108
+ loss = SamplesLoss(
109
+ loss="sinkhorn",
110
+ blur=0.1,
111
+ debias=True,
112
+ reach=1.0, # Unbalanced OT with KL penalty
113
+ )
114
+ ```
115
+
116
+ ### Semi-Unbalanced OT
117
+
118
+ Different constraints for source vs target:
119
+
120
+ ```python
121
+ loss = SamplesLoss(
122
+ loss="sinkhorn",
123
+ blur=0.1,
124
+ reach_x=1.0, # Relax source marginal
125
+ reach_y=None, # Keep target marginal strict (balanced)
126
+ )
127
+ ```
128
+
129
+ ### Early Stopping
130
+
131
+ ```python
132
+ loss = SamplesLoss(
133
+ loss="sinkhorn",
134
+ blur=0.1,
135
+ n_iters=100,
136
+ threshold=1e-3, # Stop when potential change < threshold
137
+ inner_iterations=10, # Check every N iterations
138
+ )
139
+ ```
140
+
141
+ ### Hessian-Vector Product
142
+
143
+ ```python
144
+ x = torch.randn(4096, 64, device="cuda", requires_grad=True)
145
+ y = torch.randn(4096, 64, device="cuda")
146
+ v = torch.randn_like(x)
147
+
148
+ loss = SamplesLoss(loss="sinkhorn", blur=0.1)
149
+ cost = loss(x, y)
150
+
151
+ # First-order gradient
152
+ grad_x = torch.autograd.grad(cost, x, create_graph=True)[0]
153
+
154
+ # HVP via double backward (uses streaming CG solver)
155
+ hvp = torch.autograd.grad((grad_x * v).sum(), x)[0]
156
+ ```
157
+
158
+ ## FlashSinkhorn (v0.3.0)
159
+
160
+ FlashSinkhorn is a reformulated Sinkhorn kernel that uses **shifted potentials** inspired by FlashAttention. It reduces bias vector loads by 67% and elementwise operations by 78% per tile, yielding 10-40% speedups for n >= 10,000.
161
+
162
+ ### How It Works
163
+
164
+ Standard Sinkhorn loads 3 bias vectors per tile (g, log_b, y²). FlashSinkhorn precomputes a single fused bias `u = (g_shifted + eps*log(b)) / eps` and uses raw coordinates with an inline scale factor, matching FlashAttention's score interface exactly.
165
+
166
+ ### Performance (d=64, A100-80GB, 100 iterations)
167
+
168
+ **Symmetric solver (vs v0.2.0 GeomLoss-style kernel):**
169
+
170
+ | n | v0.2.0 | v0.3.0 | Speedup |
171
+ |---|--------|--------|---------|
172
+ | 50,000 | 1730 ms | 1450 ms | **1.19x** |
173
+ | 10,000 | 88 ms | 61 ms | **1.43x** |
174
+ | 5,000 | 25 ms | 24 ms | 1.04x |
175
+
176
+ **Alternating solver (vs v0.2.0 OTT-style kernel, 10 iterations):**
177
+
178
+ | n | v0.2.0 | v0.3.0 | Speedup |
179
+ |---|--------|--------|---------|
180
+ | 50,000 | 137.9 ms | 102.6 ms | **1.34x** |
181
+ | 20,000 | 25.7 ms | 21.7 ms | **1.19x** |
182
+ | 10,000 | 8.9 ms | 8.3 ms | **1.07x** |
183
+
184
+ ### Usage
185
+
186
+ FlashSinkhorn is enabled by default (`use_flashstyle=True`):
187
+
188
+ ```python
189
+ # Default: uses FlashSinkhorn (fastest for n >= 5000)
190
+ loss = SamplesLoss(loss="sinkhorn", blur=0.1, debias=True)
191
+
192
+ # Explicitly disable to use previous kernels
193
+ loss = SamplesLoss(loss="sinkhorn", blur=0.1, debias=True, use_flashstyle=False)
194
+ ```
195
+
196
+ Low-level FlashSinkhorn API:
197
+
198
+ ```python
199
+ from flash_sinkhorn.kernels import (
200
+ sinkhorn_flashstyle_symmetric, # Full symmetric solver
201
+ sinkhorn_flashstyle_alternating, # Full alternating solver
202
+ flashsinkhorn_symmetric_step, # Single fused iteration
203
+ apply_plan_vec_flashstyle, # Transport plan @ vector (shifted potentials)
204
+ apply_plan_mat_flashstyle, # Transport plan @ matrix (shifted potentials)
205
+ )
206
+ ```
207
+
208
+ ## API Reference
209
+
210
+ ### SamplesLoss
211
+
212
+ ```python
213
+ SamplesLoss(
214
+ loss="sinkhorn",
215
+ p=2, # Only p=2 supported (squared Euclidean)
216
+ blur=0.05, # Regularization: eps = blur^2
217
+ debias=True, # Debiased Sinkhorn divergence
218
+ half_cost=False, # Use ||x-y||²/2 to match GeomLoss
219
+ reach=None, # Unbalanced OT (None = balanced)
220
+ reach_x=None, # Semi-unbalanced: source marginal
221
+ reach_y=None, # Semi-unbalanced: target marginal
222
+ scaling=0.5, # Epsilon annealing factor
223
+ n_iters=None, # Max iterations (None = use scaling)
224
+ threshold=None, # Early stopping threshold
225
+ inner_iterations=10, # Check convergence every N iters
226
+ use_flashstyle=True, # Use FlashSinkhorn shifted-potential kernels
227
+ )
228
+ ```
229
+
230
+ ### Low-Level API
231
+
232
+ ```python
233
+ # FlashSinkhorn (recommended)
234
+ from flash_sinkhorn.kernels import (
235
+ sinkhorn_flashstyle_symmetric,
236
+ sinkhorn_flashstyle_alternating,
237
+ apply_plan_vec_flashstyle,
238
+ apply_plan_mat_flashstyle,
239
+ )
240
+
241
+ # Legacy kernels (still available)
242
+ from flash_sinkhorn.kernels.sinkhorn_triton_geomloss_sqeuclid import (
243
+ sinkhorn_geomloss_online_potentials_sqeuclid,
244
+ )
245
+ from flash_sinkhorn.kernels.sinkhorn_triton_grad_sqeuclid import (
246
+ sinkhorn_geomloss_online_grad_sqeuclid,
247
+ )
248
+ from flash_sinkhorn.hvp import hvp_x_sqeuclid_from_potentials
249
+ ```
250
+
251
+ ## Key Concepts
252
+
253
+ ### Cost Convention
254
+
255
+ - **FlashSinkhorn default**: `C(x,y) = ||x-y||²`
256
+ - **GeomLoss p=2 default**: `C(x,y) = ||x-y||²/2`
257
+ - Use `half_cost=True` to match GeomLoss
258
+
259
+ ### Memory Efficiency
260
+
261
+ FlashSinkhorn streams tiles of (x,y) and computes costs on-the-fly:
262
+ - **Forward**: O(nd) memory (no n×m cost matrix)
263
+ - **Gradient**: O(nd) memory (streaming accumulation)
264
+ - **HVP**: O(nd) memory (CG solver with streaming matvec)
265
+
266
+ ### Numerical Stability
267
+
268
+ - Uses `exp2/log2` for stable LSE computation
269
+ - Safe log/division guards against underflow
270
+ - TF32 enabled by default for ~2x speedup on A100/H100 (set `allow_tf32=False` for strict FP32)
271
+ - HVP (double backward) uses strict FP32 internally for numerical stability
272
+
273
+ ## Benchmarks
274
+
275
+ Compare FlashSinkhorn against GeomLoss (KeOps) and OTT-JAX.
276
+
277
+ **Install benchmark dependencies:**
278
+ ```bash
279
+ pip install geomloss pykeops ott-jax jax[cuda12]
280
+ ```
281
+
282
+ **Run benchmarks:**
283
+ ```bash
284
+ # Forward pass benchmark
285
+ python -m flash_sinkhorn.bench.bench_forward --sizes 5000,10000,20000 --dims 64 --verify
286
+
287
+ # Backward pass benchmark
288
+ python -m flash_sinkhorn.bench.bench_backward --sizes 5000,10000,20000 --dims 64 --verify
289
+
290
+ # Quick test (small size)
291
+ python -m flash_sinkhorn.bench.bench_forward --sizes 5000 --dims 4 --verify
292
+
293
+ # Run only FlashSinkhorn (skip GeomLoss/OTT-JAX)
294
+ python -m flash_sinkhorn.bench.bench_forward --sizes 10000 --dims 64 --no-geomloss --no-ott
295
+ ```
296
+
297
+ Results are saved to `output/paper_benchmarks/forward/` and `output/paper_benchmarks/backward/`.
298
+
299
+ ## Citation
300
+
301
+ If you find FlashSinkhorn useful in your research, please cite our paper:
302
+
303
+ ```bibtex
304
+ @article{ye2026flashsinkhorn,
305
+ title={FlashSinkhorn: IO-Aware Entropic Optimal Transport},
306
+ author={Ye, Felix X.-F. and Li, Xingjie and Yu, An and Chang, Ming-Ching and Chu, Linsong and Wertheimer, Davis},
307
+ journal={arXiv preprint arXiv:2602.03067},
308
+ year={2026},
309
+ url={https://arxiv.org/abs/2602.03067}
310
+ }
311
+ ```
312
+
313
+ ## License
314
+
315
+ MIT
@@ -0,0 +1,283 @@
1
+ <p align="center">
2
+ <img src="FlashSinkhorn.png" alt="FlashSinkhorn" width="100%">
3
+ </p>
4
+
5
+ # FlashSinkhorn
6
+
7
+ **Streaming Entropic Optimal Transport in PyTorch + Triton**
8
+
9
+ FlashSinkhorn computes Sinkhorn OT using FlashAttention-style streaming—**never materializing the n×m cost matrix**—enabling **O(nd) memory** instead of O(n²).
10
+
11
+ ## Features
12
+
13
+ - **FlashSinkhorn kernels** — shifted-potential formulation inspired by FlashAttention, 10-40% faster than previous Triton kernels at n >= 10k
14
+ - **Fused Triton kernels** for forward, gradient, and HVP
15
+ - **GeomLoss-compatible API** (`SamplesLoss`)
16
+ - **Analytic gradients** (no backprop through Sinkhorn iterations)
17
+ - **Hessian-vector products** via streaming CG solver
18
+ - **Half-cost support** (`half_cost=True`) for exact GeomLoss parity
19
+ - **Unbalanced/semi-unbalanced OT** via `reach` parameter
20
+ - **Large-D support** (d > 1024) with tiled gradient kernel
21
+ - **Early stopping** with convergence threshold
22
+
23
+ ## Install
24
+
25
+ ```bash
26
+ pip install -e .
27
+ pip install -e ".[dev]" # with dev dependencies
28
+ ```
29
+
30
+ **Requirements:** PyTorch ≥2.5, Triton ≥3.1, CUDA 12.x
31
+
32
+ ## Quick Start
33
+
34
+ ### Basic Usage
35
+
36
+ ```python
37
+ import torch
38
+ from flash_sinkhorn import SamplesLoss
39
+
40
+ x = torch.randn(4096, 64, device="cuda")
41
+ y = torch.randn(4096, 64, device="cuda")
42
+
43
+ # FlashSinkhorn is the default backend (use_flashstyle=True)
44
+ loss = SamplesLoss(loss="sinkhorn", blur=0.1, debias=True)
45
+ cost = loss(x, y)
46
+ ```
47
+
48
+ ### Gradient Flow
49
+
50
+ ```python
51
+ x = torch.randn(4096, 64, device="cuda", requires_grad=True)
52
+ y = torch.randn(4096, 64, device="cuda")
53
+
54
+ loss = SamplesLoss(loss="sinkhorn", blur=0.1, debias=True)
55
+ cost = loss(x, y)
56
+ grad_x = torch.autograd.grad(cost, x)[0] # Analytic gradient
57
+ ```
58
+
59
+ ### GeomLoss Parity
60
+
61
+ Use `half_cost=True` to match GeomLoss's cost convention:
62
+
63
+ ```python
64
+ # FlashSinkhorn with half_cost matches GeomLoss exactly
65
+ flash_loss = SamplesLoss(loss="sinkhorn", blur=0.1, half_cost=True, debias=True)
66
+
67
+ # Equivalent GeomLoss call
68
+ # geomloss_loss = geomloss.SamplesLoss(loss="sinkhorn", p=2, blur=0.1, debias="positive")
69
+ ```
70
+
71
+ ### Unbalanced OT
72
+
73
+ For distributions with different total mass or outliers:
74
+
75
+ ```python
76
+ loss = SamplesLoss(
77
+ loss="sinkhorn",
78
+ blur=0.1,
79
+ debias=True,
80
+ reach=1.0, # Unbalanced OT with KL penalty
81
+ )
82
+ ```
83
+
84
+ ### Semi-Unbalanced OT
85
+
86
+ Different constraints for source vs target:
87
+
88
+ ```python
89
+ loss = SamplesLoss(
90
+ loss="sinkhorn",
91
+ blur=0.1,
92
+ reach_x=1.0, # Relax source marginal
93
+ reach_y=None, # Keep target marginal strict (balanced)
94
+ )
95
+ ```
96
+
97
+ ### Early Stopping
98
+
99
+ ```python
100
+ loss = SamplesLoss(
101
+ loss="sinkhorn",
102
+ blur=0.1,
103
+ n_iters=100,
104
+ threshold=1e-3, # Stop when potential change < threshold
105
+ inner_iterations=10, # Check every N iterations
106
+ )
107
+ ```
108
+
109
+ ### Hessian-Vector Product
110
+
111
+ ```python
112
+ x = torch.randn(4096, 64, device="cuda", requires_grad=True)
113
+ y = torch.randn(4096, 64, device="cuda")
114
+ v = torch.randn_like(x)
115
+
116
+ loss = SamplesLoss(loss="sinkhorn", blur=0.1)
117
+ cost = loss(x, y)
118
+
119
+ # First-order gradient
120
+ grad_x = torch.autograd.grad(cost, x, create_graph=True)[0]
121
+
122
+ # HVP via double backward (uses streaming CG solver)
123
+ hvp = torch.autograd.grad((grad_x * v).sum(), x)[0]
124
+ ```
125
+
126
+ ## FlashSinkhorn (v0.3.0)
127
+
128
+ FlashSinkhorn is a reformulated Sinkhorn kernel that uses **shifted potentials** inspired by FlashAttention. It reduces bias vector loads by 67% and elementwise operations by 78% per tile, yielding 10-40% speedups for n >= 10,000.
129
+
130
+ ### How It Works
131
+
132
+ Standard Sinkhorn loads 3 bias vectors per tile (g, log_b, y²). FlashSinkhorn precomputes a single fused bias `u = (g_shifted + eps*log(b)) / eps` and uses raw coordinates with an inline scale factor, matching FlashAttention's score interface exactly.
133
+
134
+ ### Performance (d=64, A100-80GB, 100 iterations)
135
+
136
+ **Symmetric solver (vs v0.2.0 GeomLoss-style kernel):**
137
+
138
+ | n | v0.2.0 | v0.3.0 | Speedup |
139
+ |---|--------|--------|---------|
140
+ | 50,000 | 1730 ms | 1450 ms | **1.19x** |
141
+ | 10,000 | 88 ms | 61 ms | **1.43x** |
142
+ | 5,000 | 25 ms | 24 ms | 1.04x |
143
+
144
+ **Alternating solver (vs v0.2.0 OTT-style kernel, 10 iterations):**
145
+
146
+ | n | v0.2.0 | v0.3.0 | Speedup |
147
+ |---|--------|--------|---------|
148
+ | 50,000 | 137.9 ms | 102.6 ms | **1.34x** |
149
+ | 20,000 | 25.7 ms | 21.7 ms | **1.19x** |
150
+ | 10,000 | 8.9 ms | 8.3 ms | **1.07x** |
151
+
152
+ ### Usage
153
+
154
+ FlashSinkhorn is enabled by default (`use_flashstyle=True`):
155
+
156
+ ```python
157
+ # Default: uses FlashSinkhorn (fastest for n >= 5000)
158
+ loss = SamplesLoss(loss="sinkhorn", blur=0.1, debias=True)
159
+
160
+ # Explicitly disable to use previous kernels
161
+ loss = SamplesLoss(loss="sinkhorn", blur=0.1, debias=True, use_flashstyle=False)
162
+ ```
163
+
164
+ Low-level FlashSinkhorn API:
165
+
166
+ ```python
167
+ from flash_sinkhorn.kernels import (
168
+ sinkhorn_flashstyle_symmetric, # Full symmetric solver
169
+ sinkhorn_flashstyle_alternating, # Full alternating solver
170
+ flashsinkhorn_symmetric_step, # Single fused iteration
171
+ apply_plan_vec_flashstyle, # Transport plan @ vector (shifted potentials)
172
+ apply_plan_mat_flashstyle, # Transport plan @ matrix (shifted potentials)
173
+ )
174
+ ```
175
+
176
+ ## API Reference
177
+
178
+ ### SamplesLoss
179
+
180
+ ```python
181
+ SamplesLoss(
182
+ loss="sinkhorn",
183
+ p=2, # Only p=2 supported (squared Euclidean)
184
+ blur=0.05, # Regularization: eps = blur^2
185
+ debias=True, # Debiased Sinkhorn divergence
186
+ half_cost=False, # Use ||x-y||²/2 to match GeomLoss
187
+ reach=None, # Unbalanced OT (None = balanced)
188
+ reach_x=None, # Semi-unbalanced: source marginal
189
+ reach_y=None, # Semi-unbalanced: target marginal
190
+ scaling=0.5, # Epsilon annealing factor
191
+ n_iters=None, # Max iterations (None = use scaling)
192
+ threshold=None, # Early stopping threshold
193
+ inner_iterations=10, # Check convergence every N iters
194
+ use_flashstyle=True, # Use FlashSinkhorn shifted-potential kernels
195
+ )
196
+ ```
197
+
198
+ ### Low-Level API
199
+
200
+ ```python
201
+ # FlashSinkhorn (recommended)
202
+ from flash_sinkhorn.kernels import (
203
+ sinkhorn_flashstyle_symmetric,
204
+ sinkhorn_flashstyle_alternating,
205
+ apply_plan_vec_flashstyle,
206
+ apply_plan_mat_flashstyle,
207
+ )
208
+
209
+ # Legacy kernels (still available)
210
+ from flash_sinkhorn.kernels.sinkhorn_triton_geomloss_sqeuclid import (
211
+ sinkhorn_geomloss_online_potentials_sqeuclid,
212
+ )
213
+ from flash_sinkhorn.kernels.sinkhorn_triton_grad_sqeuclid import (
214
+ sinkhorn_geomloss_online_grad_sqeuclid,
215
+ )
216
+ from flash_sinkhorn.hvp import hvp_x_sqeuclid_from_potentials
217
+ ```
218
+
219
+ ## Key Concepts
220
+
221
+ ### Cost Convention
222
+
223
+ - **FlashSinkhorn default**: `C(x,y) = ||x-y||²`
224
+ - **GeomLoss p=2 default**: `C(x,y) = ||x-y||²/2`
225
+ - Use `half_cost=True` to match GeomLoss
226
+
227
+ ### Memory Efficiency
228
+
229
+ FlashSinkhorn streams tiles of (x,y) and computes costs on-the-fly:
230
+ - **Forward**: O(nd) memory (no n×m cost matrix)
231
+ - **Gradient**: O(nd) memory (streaming accumulation)
232
+ - **HVP**: O(nd) memory (CG solver with streaming matvec)
233
+
234
+ ### Numerical Stability
235
+
236
+ - Uses `exp2/log2` for stable LSE computation
237
+ - Safe log/division guards against underflow
238
+ - TF32 enabled by default for ~2x speedup on A100/H100 (set `allow_tf32=False` for strict FP32)
239
+ - HVP (double backward) uses strict FP32 internally for numerical stability
240
+
241
+ ## Benchmarks
242
+
243
+ Compare FlashSinkhorn against GeomLoss (KeOps) and OTT-JAX.
244
+
245
+ **Install benchmark dependencies:**
246
+ ```bash
247
+ pip install geomloss pykeops ott-jax jax[cuda12]
248
+ ```
249
+
250
+ **Run benchmarks:**
251
+ ```bash
252
+ # Forward pass benchmark
253
+ python -m flash_sinkhorn.bench.bench_forward --sizes 5000,10000,20000 --dims 64 --verify
254
+
255
+ # Backward pass benchmark
256
+ python -m flash_sinkhorn.bench.bench_backward --sizes 5000,10000,20000 --dims 64 --verify
257
+
258
+ # Quick test (small size)
259
+ python -m flash_sinkhorn.bench.bench_forward --sizes 5000 --dims 4 --verify
260
+
261
+ # Run only FlashSinkhorn (skip GeomLoss/OTT-JAX)
262
+ python -m flash_sinkhorn.bench.bench_forward --sizes 10000 --dims 64 --no-geomloss --no-ott
263
+ ```
264
+
265
+ Results are saved to `output/paper_benchmarks/forward/` and `output/paper_benchmarks/backward/`.
266
+
267
+ ## Citation
268
+
269
+ If you find FlashSinkhorn useful in your research, please cite our paper:
270
+
271
+ ```bibtex
272
+ @article{ye2026flashsinkhorn,
273
+ title={FlashSinkhorn: IO-Aware Entropic Optimal Transport},
274
+ author={Ye, Felix X.-F. and Li, Xingjie and Yu, An and Chang, Ming-Ching and Chu, Linsong and Wertheimer, Davis},
275
+ journal={arXiv preprint arXiv:2602.03067},
276
+ year={2026},
277
+ url={https://arxiv.org/abs/2602.03067}
278
+ }
279
+ ```
280
+
281
+ ## License
282
+
283
+ MIT
@@ -0,0 +1,57 @@
1
+ [build-system]
2
+ requires = ["setuptools>=68", "wheel"]
3
+ build-backend = "setuptools.build_meta"
4
+
5
+ [project]
6
+ name = "flash-sinkhorn"
7
+ dynamic = ["version"]
8
+ description = "Sinkhorn optimal transport kernels in PyTorch + Triton (squared Euclidean, no cost matrix materialization)."
9
+ readme = "README.md"
10
+ license = "MIT"
11
+ requires-python = ">=3.9"
12
+ keywords = ["optimal-transport", "sinkhorn", "triton", "pytorch", "gpu"]
13
+ authors = [
14
+ {name = "OT Triton Contributors"}
15
+ ]
16
+ classifiers = [
17
+ "Development Status :: 3 - Alpha",
18
+ "Programming Language :: Python :: 3",
19
+ "Programming Language :: Python :: 3.9",
20
+ "Programming Language :: Python :: 3.10",
21
+ "Programming Language :: Python :: 3.11",
22
+ "Programming Language :: Python :: 3.12",
23
+ "Topic :: Scientific/Engineering",
24
+ "Environment :: GPU :: NVIDIA CUDA",
25
+ ]
26
+ dependencies = [
27
+ "numpy",
28
+ "torch>=2.5",
29
+ "triton>=3.1",
30
+ ]
31
+
32
+ [project.urls]
33
+ Homepage = "https://github.com/ot-triton-lab/flash-sinkhorn"
34
+ Repository = "https://github.com/ot-triton-lab/flash-sinkhorn"
35
+
36
+ [project.optional-dependencies]
37
+ dev = [
38
+ "pytest",
39
+ "geomloss",
40
+ "pykeops",
41
+ "jax",
42
+ "ott-jax",
43
+ "matplotlib",
44
+ ]
45
+
46
+ [tool.setuptools]
47
+ include-package-data = true
48
+
49
+ [tool.setuptools.packages.find]
50
+ where = ["src"]
51
+
52
+ [tool.setuptools.dynamic]
53
+ version = {attr = "flash_sinkhorn.__version__"}
54
+
55
+ [tool.pytest.ini_options]
56
+ addopts = "-ra"
57
+ testpaths = ["src/flash_sinkhorn/testing"]
@@ -0,0 +1,4 @@
1
+ [egg_info]
2
+ tag_build =
3
+ tag_date = 0
4
+