torchdeltaflow 0.2.2__tar.gz → 0.2.3.dev23__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 (153) hide show
  1. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/.github/workflows/publish.yml +21 -5
  2. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/CHANGELOG.md +20 -0
  3. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/PKG-INFO +5 -4
  4. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/__init__.py +5 -2
  5. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/core/__init__.py +6 -0
  6. torchdeltaflow-0.2.3.dev23/deltaflow/core/base_coupling.py +22 -0
  7. torchdeltaflow-0.2.3.dev23/deltaflow/core/base_equilibrium_field.py +31 -0
  8. torchdeltaflow-0.2.3.dev23/deltaflow/core/base_equilibrium_interpolant.py +41 -0
  9. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/interpolants/__init__.py +4 -0
  10. torchdeltaflow-0.2.3.dev23/deltaflow/interpolants/equilibrium.py +121 -0
  11. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/interpolants/linear.py +1 -1
  12. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/interpolants/ot.py +11 -42
  13. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/inverse/operators.py +2 -1
  14. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/losses/__init__.py +2 -0
  15. torchdeltaflow-0.2.3.dev23/deltaflow/losses/equilibrium_matching.py +97 -0
  16. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/models/__init__.py +9 -1
  17. torchdeltaflow-0.2.3.dev23/deltaflow/models/dit.py +438 -0
  18. torchdeltaflow-0.2.3.dev23/deltaflow/samplers/__init__.py +15 -0
  19. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/solvers/__init__.py +9 -1
  20. torchdeltaflow-0.2.3.dev23/deltaflow/solvers/gradient_descent.py +129 -0
  21. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/trainer/coupling.py +6 -13
  22. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/utils/__init__.py +2 -0
  23. torchdeltaflow-0.2.3.dev23/deltaflow/utils/ot.py +52 -0
  24. torchdeltaflow-0.2.3.dev23/docs/assets/equilibrium-matching/energy_landscape.png +0 -0
  25. torchdeltaflow-0.2.3.dev23/docs/assets/equilibrium-matching/eqm_sampling.gif +0 -0
  26. torchdeltaflow-0.2.3.dev23/docs/assets/equilibrium-matching/gd_snapshots.png +0 -0
  27. torchdeltaflow-0.2.3.dev23/docs/assets/equilibrium-matching/gd_trajectories.png +0 -0
  28. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/examples.md +80 -0
  29. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/getting-started/installation.md +6 -2
  30. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/guides/training.md +4 -3
  31. torchdeltaflow-0.2.3.dev23/examples/10-sampling/02-equilibrium-matching/main.py +72 -0
  32. torchdeltaflow-0.2.3.dev23/examples/20-training/03-conditional-dit/main.py +72 -0
  33. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/examples/90-showcase/03-minibatch-ot-viz/main.py +2 -2
  34. torchdeltaflow-0.2.3.dev23/examples/90-showcase/09-equilibrium-matching-viz/main.py +371 -0
  35. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/pyproject.toml +14 -5
  36. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/tests/test_coupling.py +42 -0
  37. torchdeltaflow-0.2.3.dev23/tests/test_dit.py +147 -0
  38. torchdeltaflow-0.2.3.dev23/tests/test_equilibrium.py +219 -0
  39. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/torchdeltaflow.egg-info/PKG-INFO +5 -4
  40. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/torchdeltaflow.egg-info/SOURCES.txt +17 -0
  41. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/torchdeltaflow.egg-info/requires.txt +1 -0
  42. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/torchdeltaflow.egg-info/scm_file_list.json +17 -0
  43. torchdeltaflow-0.2.3.dev23/torchdeltaflow.egg-info/scm_version.json +8 -0
  44. torchdeltaflow-0.2.2/deltaflow/samplers/__init__.py +0 -7
  45. torchdeltaflow-0.2.2/torchdeltaflow.egg-info/scm_version.json +0 -8
  46. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/.github/workflows/ci.yml +0 -0
  47. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/.github/workflows/docs.yml +0 -0
  48. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/.gitignore +0 -0
  49. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/CITATION.cff +0 -0
  50. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/CODE_OF_CONDUCT.md +0 -0
  51. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/CONTRIBUTING.md +0 -0
  52. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/LICENSE +0 -0
  53. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/MANIFEST.in +0 -0
  54. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/README.md +3 -3
  55. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/benchmarks/README.md +0 -0
  56. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/core/base.py +0 -0
  57. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/core/base_interpolant.py +0 -0
  58. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/core/base_loss.py +0 -0
  59. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/core/base_solver.py +0 -0
  60. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/core/base_velocity_field.py +0 -0
  61. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/datasets/__init__.py +0 -0
  62. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/datasets/radiograph.py +0 -0
  63. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/interpolants/base.py +0 -0
  64. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/interpolants/schrodinger_bridge.py +0 -0
  65. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/interpolants/variance_preserving.py +0 -0
  66. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/inverse/__init__.py +0 -0
  67. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/inverse/likelihood.py +0 -0
  68. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/inverse/tweedie.py +0 -0
  69. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/losses/conditional_flow_matching.py +0 -0
  70. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/losses/delta_alignment.py +0 -0
  71. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/losses/flow_matching.py +0 -0
  72. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/models/backbone.py +0 -0
  73. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/models/ema.py +0 -0
  74. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/models/projector.py +0 -0
  75. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/samplers/euler.py +0 -0
  76. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/solvers/euler.py +0 -0
  77. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/solvers/heun.py +0 -0
  78. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/solvers/posterior_solver.py +0 -0
  79. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/trainer/__init__.py +0 -0
  80. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/trainer/data.py +0 -0
  81. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/trainer/loop.py +0 -0
  82. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/utils/numerical.py +0 -0
  83. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/api/core.md +0 -0
  84. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/api/datasets.md +0 -0
  85. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/api/index.md +0 -0
  86. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/api/interpolants.md +0 -0
  87. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/api/inverse.md +0 -0
  88. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/api/losses.md +0 -0
  89. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/api/models.md +0 -0
  90. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/api/samplers.md +0 -0
  91. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/api/solvers.md +0 -0
  92. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/api/trainer.md +0 -0
  93. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/api/utils.md +0 -0
  94. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/algorithm-comparison/comparison.gif +0 -0
  95. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/algorithm-comparison/comparison.png +0 -0
  96. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/algorithm-comparison/trajectories_comparison.png +0 -0
  97. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/css/extra.css +0 -0
  98. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/favicon.ico +0 -0
  99. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/favicon.png +0 -0
  100. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/guidance-alignment/features.png +0 -0
  101. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/inverse-posterior/inverse_posterior.gif +0 -0
  102. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/inverse-posterior/inverse_posterior.png +0 -0
  103. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/js/mathjax.js +0 -0
  104. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/landmark-detection/landmark_detection.png +0 -0
  105. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/logo.svg +0 -0
  106. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/minibatch-ot/minibatch_ot.gif +0 -0
  107. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/minibatch-ot/minibatch_ot.png +0 -0
  108. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/sampling-flow/flow.gif +0 -0
  109. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/sampling-flow/snapshots.png +0 -0
  110. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/sampling-flow/trajectories.png +0 -0
  111. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/sampling-flow/velocity_field.png +0 -0
  112. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/schrodinger-bridge/bridge_paths.gif +0 -0
  113. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/schrodinger-bridge/bridge_paths.png +0 -0
  114. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/schrodinger-bridge/sb_flow.gif +0 -0
  115. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/schrodinger-bridge/sb_snapshots.png +0 -0
  116. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/schrodinger-bridge/sb_trajectories.png +0 -0
  117. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/concepts/architecture.md +0 -0
  118. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/contributing.md +0 -0
  119. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/getting-started/quickstart.md +0 -0
  120. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/guides/delta-alignment.md +0 -0
  121. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/guides/flow-matching.md +0 -0
  122. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/guides/inverse-problems.md +0 -0
  123. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/index.md +0 -0
  124. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/examples/00-foundations/01-linear-interpolant/main.py +0 -0
  125. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/examples/10-sampling/01-euler-flow/main.py +0 -0
  126. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/examples/20-training/01-flow-matching/main.py +0 -0
  127. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/examples/20-training/02-delta-alignment/main.py +0 -0
  128. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/examples/30-inverse/01-posterior/main.py +0 -0
  129. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/examples/90-showcase/01-landmark-viz/main.py +0 -0
  130. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/examples/90-showcase/02-sampling-flow-viz/main.py +0 -0
  131. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/examples/90-showcase/04-inverse-posterior-viz/main.py +0 -0
  132. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/examples/90-showcase/05-schrodinger-bridge-viz/main.py +0 -0
  133. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/examples/90-showcase/06-algorithm-comparison/main.py +0 -0
  134. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/examples/90-showcase/07-landmark-detection/main.py +0 -0
  135. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/examples/90-showcase/08-guidance-alignment-pretraining/main.py +0 -0
  136. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/examples/90-showcase/README.md +0 -0
  137. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/mkdocs.yml +0 -0
  138. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/overrides/main.html +0 -0
  139. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/requirements-docs.txt +0 -0
  140. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/requirements.txt +0 -0
  141. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/setup.cfg +0 -0
  142. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/tests/__init__.py +0 -0
  143. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/tests/conftest.py +0 -0
  144. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/tests/test_interpolants.py +0 -0
  145. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/tests/test_interpolants_new.py +0 -0
  146. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/tests/test_inverse.py +0 -0
  147. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/tests/test_losses.py +0 -0
  148. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/tests/test_models.py +0 -0
  149. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/tests/test_samplers.py +0 -0
  150. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/tests/test_solvers.py +0 -0
  151. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/tests/test_trainer.py +0 -0
  152. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/torchdeltaflow.egg-info/dependency_links.txt +0 -0
  153. {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/torchdeltaflow.egg-info/top_level.txt +0 -0
@@ -1,11 +1,18 @@
1
1
  name: publish
2
2
 
3
- # Publishes torchdeltaflow to PyPI whenever a version tag (e.g. v0.2.2) is
4
- # pushed. The version is derived from the tag by setuptools_scm, so the tag is
5
- # the single source of truth. To cut a release, run:
3
+ # Publishes torchdeltaflow to PyPI in two cases:
6
4
  #
7
- # git tag v0.2.2
8
- # git push origin v0.2.2
5
+ # 1. A version tag (e.g. v0.2.2) is pushed -> a stable release. The version is
6
+ # derived from the tag by setuptools_scm, so the tag is the single source
7
+ # of truth. To cut a release, run:
8
+ #
9
+ # git tag v0.2.2
10
+ # git push origin v0.2.2
11
+ #
12
+ # 2. A pull request is merged into main -> a dev pre-release. setuptools_scm
13
+ # derives a clean PEP 440 version from the commit distance to the last tag
14
+ # (e.g. "0.2.3.dev4"), which is unique per merge. "skip-existing" below
15
+ # guards against the rare case where the same version was already uploaded.
9
16
  #
10
17
  # Authentication uses a PyPI API token stored as the repository secret
11
18
  # PYPI_API_TOKEN. Create the token at https://pypi.org/manage/account/token/
@@ -16,12 +23,18 @@ on:
16
23
  push:
17
24
  tags:
18
25
  - "v*"
26
+ pull_request:
27
+ branches: [main]
28
+ types: [closed]
19
29
 
20
30
  permissions:
21
31
  contents: read
22
32
 
23
33
  jobs:
24
34
  build:
35
+ # Run for pushed version tags, or when a PR is actually merged (not just
36
+ # closed). Plain "closed" without merge must not publish.
37
+ if: github.event_name == 'push' || github.event.pull_request.merged == true
25
38
  runs-on: ubuntu-latest
26
39
  steps:
27
40
  - uses: actions/checkout@v4
@@ -60,3 +73,6 @@ jobs:
60
73
  uses: pypa/gh-action-pypi-publish@release/v1
61
74
  with:
62
75
  password: ${{ secrets.PYPI_API_TOKEN }}
76
+ # Merge builds may reproduce an existing dev version if no new commits
77
+ # changed the distance to the last tag. Skip rather than fail.
78
+ skip-existing: true
@@ -7,6 +7,26 @@ The format follows [Keep a Changelog](https://keepachangelog.com/en/1.1.0/).
7
7
 
8
8
  ### Added
9
9
 
10
+ - `deltaflow.models.DiT`: a class-conditional Diffusion Transformer velocity
11
+ field with **adaLN-Zero** conditioning (Peebles & Xie, 2023). Ships a reusable
12
+ `TimestepEmbedding`, a `LabelEmbedding` with a learned null token for
13
+ classifier-free guidance, zero-initialized modulation gates (identity at init,
14
+ zero initial velocity), `cond["y"]` support through the existing
15
+ `BaseVelocityField` API, and a `forward_with_cfg` guided-sampling helper. See
16
+ the new `examples/20-training/03-conditional-dit` walkthrough.
17
+
18
+ ### Changed
19
+
20
+ - `scipy` is now a core dependency (previously gated behind the `ot`
21
+ extra), so `OTInterpolant`/`OTCoupling` use the exact Hungarian-algorithm
22
+ mini-batch OT assignment by default. The `ot` extra is kept as a
23
+ backwards-compatible no-op, and the greedy nearest-neighbour matcher
24
+ remains as a defensive fallback if `scipy` is ever missing.
25
+
26
+ ## [0.2.2]
27
+
28
+ ### Added
29
+
10
30
  - API-reference docstrings now include the underlying mathematics in
11
31
  LaTeX (rendered via MathJax/arithmatex): probability paths and target
12
32
  velocities for the interpolants (`Linear`, `VariancePreserving`, `OT`,
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: torchdeltaflow
3
- Version: 0.2.2
3
+ Version: 0.2.3.dev23
4
4
  Summary: Flow matching, mini-batch optimal transport coupling, and posterior sampling for inverse problems in PyTorch
5
5
  Author-email: Phrugsa Limbunlom <phrugsa.lim@gmail.com>
6
6
  Maintainer-email: Phrugsa Limbunlom <phrugsa.lim@gmail.com>
@@ -53,6 +53,7 @@ Requires-Dist: torch>=2.0
53
53
  Requires-Dist: numpy
54
54
  Requires-Dist: einops
55
55
  Requires-Dist: tqdm
56
+ Requires-Dist: scipy>=1.10
56
57
  Provides-Extra: images
57
58
  Requires-Dist: pillow>=9.0; extra == "images"
58
59
  Provides-Extra: ot
@@ -85,6 +86,9 @@ Dynamic: license-file
85
86
  <a href="https://github.com/phrugsa-limbunlom/deltaflow/blob/main/LICENSE" target="_blank" title="License">
86
87
  <img alt="License" src="https://img.shields.io/github/license/phrugsa-limbunlom/deltaflow?style=flat-square&color=3f9e73">
87
88
  </a>
89
+ <a href="https://pepy.tech/project/torchdeltaflow" target="_blank" title="Downloads">
90
+ <img alt="Downloads" src="https://static.pepy.tech/badge/torchdeltaflow?style=flat-square">
91
+ </a>
88
92
  <a href="https://github.com/phrugsa-limbunlom/deltaflow" target="_blank" title="GitHub Repo Stars">
89
93
  <img alt="GitHub Stars" src="https://img.shields.io/github/stars/phrugsa-limbunlom/deltaflow?style=social">
90
94
  </a>
@@ -97,9 +101,6 @@ Dynamic: license-file
97
101
  <a href="https://github.com/phrugsa-limbunlom/deltaflow/actions/workflows/docs.yml" target="_blank" title="Documentation">
98
102
  <img alt="Docs" src="https://img.shields.io/github/actions/workflow/status/phrugsa-limbunlom/deltaflow/docs.yml?branch=main&style=flat-square&label=docs&color=0f9bab">
99
103
  </a>
100
- <a href="https://pepy.tech/project/torchdeltaflow" target="_blank" title="Downloads">
101
- <img alt="Downloads" src="https://static.pepy.tech/badge/torchdeltaflow?style=flat-square">
102
- </a>
103
104
  <a href="https://pypi.org/project/torchdeltaflow/" target="_blank" title="Python Versions">
104
105
  <img alt="Python Versions" src="https://img.shields.io/pypi/pyversions/torchdeltaflow?style=flat-square&color=e0b13c">
105
106
  </a>
@@ -8,10 +8,13 @@ Module map:
8
8
  subclasses (`BaseVelocityField`, `BaseInterpolant`,
9
9
  `BaseSolver`, `BaseLoss`).
10
10
  - `deltaflow.interpolants`, probability paths (linear, mini-batch OT,
11
- variance-preserving).
11
+ variance-preserving, Schrödinger bridge, and the Equilibrium Matching
12
+ energy-compatible target).
12
13
  - `deltaflow.losses`, conditional flow matching, plus the
13
14
  optional delta-alignment loss for guidance-representation pretraining.
14
- - `deltaflow.solvers`, Euler, Heun, and the
15
+ - `deltaflow.solvers`, Euler, Heun, the `EquilibriumSolver` which samples an
16
+ Equilibrium Matching field by gradient descent on its implicit energy
17
+ landscape, and the
15
18
  `PosteriorSolver` which *wraps* a base solver
16
19
  and injects the measurement-likelihood gradient per step (FlowDPS /
17
20
  Flower style).
@@ -5,12 +5,18 @@ subclasses one of these bases, so new variants are drop-in and not
5
5
  rewrites of the surrounding machinery.
6
6
  """
7
7
 
8
+ from .base_coupling import BaseCoupling
9
+ from .base_equilibrium_field import BaseEquilibriumField
10
+ from .base_equilibrium_interpolant import BaseEquilibriumInterpolant
8
11
  from .base_interpolant import BaseInterpolant
9
12
  from .base_loss import BaseLoss
10
13
  from .base_solver import BaseSolver
11
14
  from .base_velocity_field import BaseVelocityField
12
15
 
13
16
  __all__ = [
17
+ "BaseCoupling",
18
+ "BaseEquilibriumField",
19
+ "BaseEquilibriumInterpolant",
14
20
  "BaseInterpolant",
15
21
  "BaseLoss",
16
22
  "BaseSolver",
@@ -0,0 +1,22 @@
1
+ """Base class for train-time coupling strategies between ``x0`` and ``x1``.
2
+
3
+ Coupling is the choice of which ``(x0, x1)`` pairs the flow-matching
4
+ regression is computed on, kept deliberately separate from the choice of
5
+ probability path (`deltaflow.interpolants`). Concrete strategies live in
6
+ `deltaflow.trainer.coupling`, this base sits in ``core`` alongside the other
7
+ drop-in abstractions so a new coupling is a subclass rather than a rewrite of
8
+ the surrounding training loop.
9
+ """
10
+
11
+ from abc import ABC, abstractmethod
12
+ from typing import Tuple
13
+
14
+ import torch
15
+
16
+
17
+ class BaseCoupling(ABC):
18
+ """Given a batch ``x1`` of data samples, return a paired ``(x0, x1)``."""
19
+
20
+ @abstractmethod
21
+ def sample_pair(self, x1: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
22
+ raise NotImplementedError
@@ -0,0 +1,31 @@
1
+ """Base class for time-invariant Equilibrium Matching fields ``f(x)``."""
2
+
3
+ from abc import ABC, abstractmethod
4
+
5
+ import torch
6
+ import torch.nn as nn
7
+
8
+
9
+ class BaseEquilibriumField(nn.Module, ABC):
10
+ r"""Base class for a time-invariant Equilibrium Matching field \(f(x)\).
11
+
12
+ Unlike a flow-matching velocity field \(v_\theta(x, t)\), an EqM field
13
+ carries no time argument. It approximates the equilibrium gradient of an
14
+ implicit energy, \(f_\theta(x) \approx -\nabla E(x)\), pointing from noise
15
+ toward data. The interpolation coefficient \(\gamma\) is implicit and (per
16
+ the EqM paper) never seen by the model, so ``forward`` takes only ``x`` and
17
+ any extra conditioning, never a time or \(\gamma\).
18
+
19
+ Subclasses implement `forward` and return a tensor with the same
20
+ shape as ``x``. Additional conditioning is passed as keyword arguments and
21
+ forwarded unchanged by `EquilibriumMatchingLoss` and
22
+ `EquilibriumSolver`.
23
+
24
+ References:
25
+ Wang and Du, "Equilibrium Matching: Generative Modeling with Implicit
26
+ Energy-Based Models" (2025), https://arxiv.org/abs/2510.02300.
27
+ """
28
+
29
+ @abstractmethod
30
+ def forward(self, x: torch.Tensor, **cond) -> torch.Tensor:
31
+ raise NotImplementedError
@@ -0,0 +1,41 @@
1
+ """Base class for Equilibrium Matching interpolants (parametrised by gamma)."""
2
+
3
+ from abc import ABC, abstractmethod
4
+ from typing import Optional, Tuple
5
+
6
+ import torch
7
+
8
+
9
+ class BaseEquilibriumInterpolant(ABC):
10
+ r"""Base class for an energy-compatible path indexed by a coefficient gamma.
11
+
12
+ This base is deliberately separate from `BaseInterpolant`. A flow-matching
13
+ interpolant is parametrised by a dynamical time ``t`` that a sampler
14
+ integrates over. An Equilibrium Matching interpolant is parametrised by an
15
+ interpolation coefficient \(\gamma \in [0, 1]\), a noise level rather than a
16
+ time. The learned field is trained to be time-invariant, so \(\gamma\) only
17
+ indexes where along the noise-to-data path a training point sits, it is
18
+ never integrated and (per the EqM paper) is not seen by the model.
19
+
20
+ Convention: \(\gamma = 0\) is noise (\(x_\gamma = x_0\)) and \(\gamma = 1\)
21
+ is data (\(x_\gamma = x_1\)).
22
+
23
+ References:
24
+ Wang and Du, "Equilibrium Matching: Generative Modeling with Implicit
25
+ Energy-Based Models" (2025), https://arxiv.org/abs/2510.02300.
26
+ """
27
+
28
+ @abstractmethod
29
+ def interpolate(
30
+ self, x1: torch.Tensor, gamma: torch.Tensor, x0: Optional[torch.Tensor] = None
31
+ ) -> Tuple[torch.Tensor, torch.Tensor]:
32
+ r"""Return ``(x_gamma, target)`` for the given data and coefficient.
33
+
34
+ Args:
35
+ x1: data sample (the \(\gamma = 1\) endpoint).
36
+ gamma: interpolation coefficient (noise level) in \([0, 1]\), not a
37
+ dynamical time.
38
+ x0: optional noise sample (the \(\gamma = 0\) endpoint). Drawn from
39
+ a standard normal when omitted.
40
+ """
41
+ raise NotImplementedError
@@ -1,13 +1,17 @@
1
1
  """Probability paths connecting noise ``x0`` to data ``x1``."""
2
2
 
3
+ from ..core.base_equilibrium_interpolant import BaseEquilibriumInterpolant
3
4
  from ..core.base_interpolant import BaseInterpolant
5
+ from .equilibrium import EquilibriumInterpolant
4
6
  from .linear import LinearInterpolant
5
7
  from .ot import OTInterpolant
6
8
  from .schrodinger_bridge import SchrodingerBridgeInterpolant
7
9
  from .variance_preserving import VariancePreservingInterpolant
8
10
 
9
11
  __all__ = [
12
+ "BaseEquilibriumInterpolant",
10
13
  "BaseInterpolant",
14
+ "EquilibriumInterpolant",
11
15
  "LinearInterpolant",
12
16
  "OTInterpolant",
13
17
  "SchrodingerBridgeInterpolant",
@@ -0,0 +1,121 @@
1
+ """Equilibrium Matching probability path (energy-compatible target)."""
2
+
3
+ from typing import Optional, Tuple
4
+
5
+ import torch
6
+
7
+ from ..core.base_equilibrium_interpolant import BaseEquilibriumInterpolant
8
+
9
+
10
+ class EquilibriumInterpolant(BaseEquilibriumInterpolant):
11
+ r"""Straight-line path with an energy-compatible (equilibrium) target.
12
+
13
+ Equilibrium Matching (EqM) keeps the rectified-flow straight-line path but
14
+ parametrises it by an interpolation coefficient \(\gamma \in [0, 1]\), a
15
+ noise level rather than a dynamical time. The learned field is trained to be
16
+ *time-invariant* (it approximates an equilibrium gradient), so \(\gamma\)
17
+ only indexes where along the noise-to-data path a training point sits, it is
18
+ not integrated over. In the paper's notation, for data \(x\) and Gaussian
19
+ noise \(\epsilon\),
20
+
21
+ \[
22
+ x_\gamma = \gamma\,x + (1 - \gamma)\,\epsilon,
23
+ \]
24
+
25
+ which, with DeltaFlow's \((x_0, x_1)\) = (noise, data) naming, is the same
26
+ straight line
27
+
28
+ \[
29
+ x_\gamma = (1 - \gamma)\,x_0 + \gamma\,x_1, \qquad \gamma \in [0, 1].
30
+ \]
31
+
32
+ EqM reshapes the regression target so the learned field becomes the
33
+ gradient of an implicit energy landscape rather than a time-conditional
34
+ velocity. Instead of regressing onto the constant displacement
35
+ \(x_1 - x_0\), it regresses onto the scaled displacement (the paper's
36
+ \((x - \epsilon)\,c(\gamma)\)),
37
+
38
+ \[
39
+ u_\gamma = c(\gamma)\,(x_1 - x_0),
40
+ \]
41
+
42
+ where \(c(\gamma)\) is the equilibrium coefficient. The key design
43
+ constraint is \(c(1) = 0\): the target vanishes at data, so ground-truth
44
+ samples become stationary points (local minima) of the landscape whose
45
+ gradient the field learns. Away from data the coefficient is held on a
46
+ constant plateau, so the field points from noise toward data with roughly
47
+ constant magnitude.
48
+
49
+ Concretely the coefficient is the minimum of two lines, rescaled by
50
+ ``scale``,
51
+
52
+ \[
53
+ c(\gamma) = \text{scale}\cdot\min\!\Bigl(
54
+ \text{start} - \tfrac{\text{start} - 1}{p}\,\gamma,\;
55
+ \tfrac{1 - \gamma}{1 - p}
56
+ \Bigr),
57
+ \]
58
+
59
+ with plateau fraction \(p\) (``plateau``). With the defaults
60
+ (``start = 1``, ``plateau = 0.8``, ``scale = 4``) the first line is flat at
61
+ \(1\), so \(c(\gamma) = 4\,\min(1, 5(1 - \gamma))\): a plateau of \(4\) for
62
+ \(\gamma \le 0.8\) that then ramps linearly down to \(0\) at \(\gamma = 1\).
63
+
64
+ **Interpolation coefficient.** \(\gamma = 0\) is noise (\(x_\gamma = x_0\))
65
+ and \(\gamma = 1\) is data (\(x_\gamma = x_1\)). This interpolant subclasses
66
+ `BaseEquilibriumInterpolant`, whose ``interpolate`` takes ``gamma`` rather
67
+ than the dynamical time ``t`` of `BaseInterpolant`, because \(\gamma\) is a
68
+ noise level and not a time to be integrated.
69
+
70
+ **Coupling.** This interpolant defines only the *path* and *target*. It is
71
+ agnostic to how \((x_0, x_1)\) pairs are formed. If ``x0`` is not supplied
72
+ it is drawn from \(\mathcal{N}(0, I)\) independently of ``x1``. Pass an
73
+ `OTCoupling` on the training side for
74
+ straighter, OT-coupled displacements.
75
+
76
+ **Sampling.** A field trained with this target is *not* integrated over a
77
+ fixed time horizon. Because it approximates an equilibrium gradient it is
78
+ sampled by optimisation (gradient descent on the landscape), for which see
79
+ `EquilibriumSolver`.
80
+
81
+ References:
82
+ Wang and Du, "Equilibrium Matching: Generative Modeling with Implicit
83
+ Energy-Based Models" (2025), https://arxiv.org/abs/2510.02300.
84
+
85
+ Args:
86
+ plateau: fraction \(p \in (0, 1)\) of the path spent before the
87
+ coefficient starts ramping down to zero. The ramp occupies the
88
+ final ``1 - plateau`` of the \([0, 1]\) interval.
89
+ scale: overall magnitude applied to the coefficient. Sets the typical
90
+ norm of the learned equilibrium gradient away from data.
91
+ start: value of the first (upper) line at ``gamma=0``. Left at ``1``
92
+ this line stays flat so the coefficient plateaus at ``scale``,
93
+ values other than ``1`` tilt the plateau.
94
+ """
95
+
96
+ def __init__(self, plateau: float = 0.8, scale: float = 4.0, start: float = 1.0):
97
+ if not 0.0 < plateau < 1.0:
98
+ raise ValueError(f"plateau must lie in (0, 1), got {plateau!r}")
99
+ self.plateau = plateau
100
+ self.scale = scale
101
+ self.start = start
102
+
103
+ def equilibrium_coefficient(self, gamma: torch.Tensor) -> torch.Tensor:
104
+ r"""Return the scalar coefficient \(c(\gamma)\) applied to \(x_1 - x_0\).
105
+
106
+ ``gamma`` is the interpolation coefficient (noise level), \(0\) at noise
107
+ and \(1\) at data, not a dynamical time.
108
+ """
109
+ line_plateau = self.start - (self.start - 1.0) / self.plateau * gamma
110
+ line_ramp = (1.0 - gamma) / (1.0 - self.plateau)
111
+ return self.scale * torch.minimum(line_plateau, line_ramp)
112
+
113
+ def interpolate(
114
+ self, x1: torch.Tensor, gamma: torch.Tensor, x0: Optional[torch.Tensor] = None
115
+ ) -> Tuple[torch.Tensor, torch.Tensor]:
116
+ if x0 is None:
117
+ x0 = torch.randn_like(x1)
118
+ gamma = gamma.view(-1, *([1] * (x1.dim() - 1)))
119
+ x_gamma = (1 - gamma) * x0 + gamma * x1
120
+ target_v = self.equilibrium_coefficient(gamma) * (x1 - x0)
121
+ return x_gamma, target_v
@@ -43,7 +43,7 @@ class LinearInterpolant(BaseInterpolant):
43
43
  ) -> Tuple[torch.Tensor, torch.Tensor]:
44
44
  if x0 is None:
45
45
  x0 = torch.randn_like(x1)
46
- t_ = t.view(-1, *([1] * (x1.dim() - 1)))
46
+ t_ = t.view(-1, *([1] * (x1.dim() - 1))) # reshape t to match the dimensions of x1 for broadcasting e.g., (B, 1, 1, 1) -> (B, C, H, W) during operation
47
47
  x_t = (1 - t_) * x0 + t_ * x1
48
48
  target_v = x1 - x0
49
49
  return x_t, target_v
@@ -10,9 +10,10 @@ optimal transport" (arXiv:2302.00482), and the OT-vs-independent ablation
10
10
  in "Flower: A Flow-Matching Solver for Inverse Problems" (arXiv:2509.26287)
11
11
  for the sampling side.
12
12
 
13
- Exact optimal assignment (Hungarian algorithm) is used when ``scipy`` is
14
- installed, otherwise a deterministic greedy nearest-neighbour fallback is
15
- used. This is a coupling strategy, not a new probability path, the same
13
+ Exact optimal assignment (Hungarian algorithm) is used by default, since
14
+ ``scipy`` is a core dependency. A deterministic greedy nearest-neighbour
15
+ fallback is used only if ``scipy`` is unavailable in the environment. This
16
+ is a coupling strategy, not a new probability path, the same
16
17
  straight-line ``x_t = (1-t) x0 + t x1`` interpolation is applied after
17
18
  permuting.
18
19
  """
@@ -22,43 +23,11 @@ from typing import Optional, Tuple
22
23
  import torch
23
24
 
24
25
  from ..core.base_interpolant import BaseInterpolant
26
+ from ..utils.ot import batch_ot_permutation
25
27
  from .linear import LinearInterpolant
26
28
 
27
-
28
- def _batch_ot_permutation(x0: torch.Tensor, x1: torch.Tensor) -> torch.Tensor:
29
- """Return a permutation ``perm`` such that ``x0[perm]`` is OT-coupled to ``x1``.
30
-
31
- Costs are squared L2 distances on the flattened per-sample tensors.
32
- """
33
- b = x0.shape[0]
34
- if b == 1:
35
- return torch.zeros(1, dtype=torch.long, device=x0.device)
36
-
37
- x0f = x0.reshape(b, -1).float()
38
- x1f = x1.reshape(b, -1).float()
39
- cost = torch.cdist(x0f, x1f) ** 2 # (B, B), cost[i, j] = |x0[i] - x1[j]|^2
40
-
41
- try:
42
- from scipy.optimize import linear_sum_assignment
43
-
44
- row_ind, col_ind = linear_sum_assignment(cost.detach().cpu().numpy())
45
- # linear_sum_assignment guarantees row_ind == 0..B-1 in sorted order;
46
- # col_ind[i] is the x1 index paired with x0[i]. We want a permutation
47
- # of x0 aligned to x1's original order: for each x1[j], take x0[i] where col_ind[i] == j.
48
- col = torch.as_tensor(col_ind, dtype=torch.long, device=x0.device)
49
- perm = torch.argsort(col)
50
- return perm
51
- except ImportError:
52
- # Greedy fallback: for each x1[j] in order, pick the closest un-used x0[i].
53
- used = torch.zeros(b, dtype=torch.bool, device=x0.device)
54
- perm = torch.empty(b, dtype=torch.long, device=x0.device)
55
- for j in range(b):
56
- row = cost[:, j].clone()
57
- row[used] = float("inf")
58
- i = int(torch.argmin(row).item())
59
- perm[j] = i
60
- used[i] = True
61
- return perm
29
+ # Backward-compat alias: the implementation now lives in ``deltaflow.utils.ot``.
30
+ _batch_ot_permutation = batch_ot_permutation
62
31
 
63
32
 
64
33
  class OTInterpolant(BaseInterpolant):
@@ -91,9 +60,9 @@ class OTInterpolant(BaseInterpolant):
91
60
  objective is identical to standard conditional flow matching and no other
92
61
  component (loss, solver, model) needs to change.
93
62
 
94
- **Solver.** The exact assignment (Hungarian algorithm) is used when
95
- ``scipy`` is installed, otherwise a deterministic greedy nearest-neighbour
96
- fallback is used.
63
+ **Solver.** The exact assignment (Hungarian algorithm) is used by
64
+ default, since ``scipy`` is a core dependency. A deterministic greedy
65
+ nearest-neighbour fallback is used only if ``scipy`` is unavailable.
97
66
 
98
67
  References:
99
68
  Tong et al., "Improving and generalizing flow-based generative
@@ -110,6 +79,6 @@ class OTInterpolant(BaseInterpolant):
110
79
  ) -> Tuple[torch.Tensor, torch.Tensor]:
111
80
  if x0 is None:
112
81
  x0 = torch.randn_like(x1)
113
- perm = _batch_ot_permutation(x0, x1)
82
+ perm = batch_ot_permutation(x0, x1)
114
83
  x0 = x0[perm]
115
84
  return self._linear.interpolate(x1, t, x0=x0)
@@ -82,9 +82,10 @@ class BlurOperator(nn.Module):
82
82
 
83
83
  def forward(self, x: torch.Tensor) -> torch.Tensor:
84
84
  pad = self.kernel_size // 2
85
+ kernel: torch.Tensor = self.get_buffer("kernel")
85
86
  return F.conv2d(
86
87
  x,
87
- self.kernel.to(dtype=x.dtype, device=x.device),
88
+ kernel.to(dtype=x.dtype, device=x.device),
88
89
  padding=pad,
89
90
  groups=self.channels,
90
91
  )
@@ -2,10 +2,12 @@
2
2
 
3
3
  from .conditional_flow_matching import ConditionalFlowMatchingLoss, FlowMatchingLoss
4
4
  from .delta_alignment import DeltaAlignmentLoss, delta_alignment_loss
5
+ from .equilibrium_matching import EquilibriumMatchingLoss
5
6
 
6
7
  __all__ = [
7
8
  "ConditionalFlowMatchingLoss",
8
9
  "DeltaAlignmentLoss",
10
+ "EquilibriumMatchingLoss",
9
11
  "FlowMatchingLoss",
10
12
  "delta_alignment_loss",
11
13
  ]
@@ -0,0 +1,97 @@
1
+ """Equilibrium Matching loss (energy-compatible target, time-invariant field)."""
2
+
3
+ from typing import TYPE_CHECKING, Optional
4
+
5
+ import torch
6
+ import torch.nn.functional as F
7
+
8
+ from ..core.base_equilibrium_interpolant import BaseEquilibriumInterpolant
9
+ from ..core.base_loss import BaseLoss
10
+ from ..interpolants.equilibrium import EquilibriumInterpolant
11
+
12
+ if TYPE_CHECKING:
13
+ from ..trainer.coupling import BaseCoupling
14
+
15
+
16
+ class EquilibriumMatchingLoss(BaseLoss):
17
+ r"""Regress a time-invariant field onto the Equilibrium Matching target.
18
+
19
+ Equilibrium Matching (EqM) trains a field \(f_\theta\) to match an
20
+ energy-compatible target along the straight noise-to-data path. For an
21
+ interpolation coefficient \(\gamma \in [0, 1]\) (a noise level), noise
22
+ \(x_0\) and data \(x_1\), the corrupted sample and target are produced by
23
+ an `EquilibriumInterpolant`,
24
+
25
+ \[
26
+ x_\gamma = (1 - \gamma)\,x_0 + \gamma\,x_1, \qquad
27
+ u_\gamma = c(\gamma)\,(x_1 - x_0),
28
+ \]
29
+
30
+ and the objective is (the paper's Eq. 3, up to the target-sign convention
31
+ below)
32
+
33
+ \[
34
+ \mathcal{L}_{\text{EqM}} =
35
+ \mathbb{E}_{\gamma,\, x_0,\, x_1}
36
+ \bigl\| f_\theta(x_\gamma) - u_\gamma \bigr\|^2 .
37
+ \]
38
+
39
+ **\(\gamma\) is implicit, and there is no time at all.** Unlike flow
40
+ matching, where the model is conditioned on time \(t\), an EqM field is
41
+ time-invariant and noise-unconditional, \(f_\theta(x)\). The coefficient
42
+ \(\gamma\) is *not seen by the model*, and neither is any surrogate time.
43
+ This loss therefore queries the model as ``model(x, **cond)``, passing no
44
+ time (and no \(\gamma\)) at all, matching `BaseEquilibriumField`.
45
+
46
+ **Target sign.** The interpolant returns \(u_\gamma = c(\gamma)(x_1 - x_0)\)
47
+ (data minus noise), so the field points from noise toward data and is
48
+ sampled by gradient *ascent* \(x \leftarrow x + \eta f(x)\) in
49
+ `EquilibriumSolver`. This matches the official EqM code. The
50
+ paper writes the mirror-image \((\epsilon - x)c(\gamma)\) with descent, the
51
+ two conventions are equivalent under a global sign flip.
52
+
53
+ References:
54
+ Wang and Du, "Equilibrium Matching: Generative Modeling with Implicit
55
+ Energy-Based Models" (2025), https://arxiv.org/abs/2510.02300.
56
+
57
+ Args:
58
+ interpolant: the energy-compatible path. Defaults to
59
+ `EquilibriumInterpolant`.
60
+ coupling: optional train-time coupling that produces \((x_0, x_1)\)
61
+ pairs from a batch of \(x_1\). See
62
+ `deltaflow.trainer.coupling`.
63
+ loss_type: one of ``"l2"``, ``"l1"``, ``"huber"``.
64
+ """
65
+
66
+ def __init__(
67
+ self,
68
+ interpolant: Optional[BaseEquilibriumInterpolant] = None,
69
+ coupling: Optional["BaseCoupling"] = None,
70
+ loss_type: str = "l2",
71
+ ):
72
+ self.interpolant = interpolant or EquilibriumInterpolant()
73
+ self.coupling = coupling
74
+ self.loss_type = loss_type
75
+
76
+ def _reduce(self, target: torch.Tensor, pred: torch.Tensor) -> torch.Tensor:
77
+ if self.loss_type == "l1":
78
+ return F.l1_loss(target, pred)
79
+ if self.loss_type == "l2":
80
+ return F.mse_loss(target, pred)
81
+ if self.loss_type == "huber":
82
+ return F.smooth_l1_loss(target, pred)
83
+ raise NotImplementedError(f"Unknown loss_type: {self.loss_type!r}")
84
+
85
+ def __call__(self, model, x1: torch.Tensor, **cond) -> torch.Tensor:
86
+ if self.coupling is not None:
87
+ x0, x1 = self.coupling.sample_pair(x1)
88
+ else:
89
+ x0 = None
90
+ gamma = torch.rand(x1.shape[0], device=x1.device)
91
+ x_gamma, target = self.interpolant.interpolate(x1, gamma, x0=x0)
92
+ # gamma is implicit and there is no time: the field is queried as f(x).
93
+ pred = model(x_gamma, **cond)
94
+ return self._reduce(target, pred)
95
+
96
+
97
+ __all__ = ["EquilibriumMatchingLoss"]
@@ -1,13 +1,21 @@
1
- """Model components: velocity-field backbone wrappers, projector heads, EMA."""
1
+ """Model components: velocity-field backbone wrappers, projector heads, EMA,
2
+ and the DiT transformer with adaLN-Zero conditioning."""
2
3
 
3
4
  from .backbone import TinyVelocityField, WrappedBackbone
5
+ from .dit import DiT, DiTBlock, FinalLayer, LabelEmbedding, TimestepEmbedding, modulate
4
6
  from .ema import EMA
5
7
  from .projector import MultiScaleProjector, ProjectorHead
6
8
 
7
9
  __all__ = [
10
+ "DiT",
11
+ "DiTBlock",
8
12
  "EMA",
13
+ "FinalLayer",
14
+ "LabelEmbedding",
9
15
  "MultiScaleProjector",
10
16
  "ProjectorHead",
17
+ "TimestepEmbedding",
11
18
  "TinyVelocityField",
12
19
  "WrappedBackbone",
20
+ "modulate",
13
21
  ]