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.
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/.github/workflows/publish.yml +21 -5
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/CHANGELOG.md +20 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/PKG-INFO +5 -4
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/__init__.py +5 -2
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/core/__init__.py +6 -0
- torchdeltaflow-0.2.3.dev23/deltaflow/core/base_coupling.py +22 -0
- torchdeltaflow-0.2.3.dev23/deltaflow/core/base_equilibrium_field.py +31 -0
- torchdeltaflow-0.2.3.dev23/deltaflow/core/base_equilibrium_interpolant.py +41 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/interpolants/__init__.py +4 -0
- torchdeltaflow-0.2.3.dev23/deltaflow/interpolants/equilibrium.py +121 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/interpolants/linear.py +1 -1
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/interpolants/ot.py +11 -42
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/inverse/operators.py +2 -1
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/losses/__init__.py +2 -0
- torchdeltaflow-0.2.3.dev23/deltaflow/losses/equilibrium_matching.py +97 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/models/__init__.py +9 -1
- torchdeltaflow-0.2.3.dev23/deltaflow/models/dit.py +438 -0
- torchdeltaflow-0.2.3.dev23/deltaflow/samplers/__init__.py +15 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/solvers/__init__.py +9 -1
- torchdeltaflow-0.2.3.dev23/deltaflow/solvers/gradient_descent.py +129 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/trainer/coupling.py +6 -13
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/utils/__init__.py +2 -0
- torchdeltaflow-0.2.3.dev23/deltaflow/utils/ot.py +52 -0
- torchdeltaflow-0.2.3.dev23/docs/assets/equilibrium-matching/energy_landscape.png +0 -0
- torchdeltaflow-0.2.3.dev23/docs/assets/equilibrium-matching/eqm_sampling.gif +0 -0
- torchdeltaflow-0.2.3.dev23/docs/assets/equilibrium-matching/gd_snapshots.png +0 -0
- torchdeltaflow-0.2.3.dev23/docs/assets/equilibrium-matching/gd_trajectories.png +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/examples.md +80 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/getting-started/installation.md +6 -2
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/guides/training.md +4 -3
- torchdeltaflow-0.2.3.dev23/examples/10-sampling/02-equilibrium-matching/main.py +72 -0
- torchdeltaflow-0.2.3.dev23/examples/20-training/03-conditional-dit/main.py +72 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/examples/90-showcase/03-minibatch-ot-viz/main.py +2 -2
- torchdeltaflow-0.2.3.dev23/examples/90-showcase/09-equilibrium-matching-viz/main.py +371 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/pyproject.toml +14 -5
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/tests/test_coupling.py +42 -0
- torchdeltaflow-0.2.3.dev23/tests/test_dit.py +147 -0
- torchdeltaflow-0.2.3.dev23/tests/test_equilibrium.py +219 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/torchdeltaflow.egg-info/PKG-INFO +5 -4
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/torchdeltaflow.egg-info/SOURCES.txt +17 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/torchdeltaflow.egg-info/requires.txt +1 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/torchdeltaflow.egg-info/scm_file_list.json +17 -0
- torchdeltaflow-0.2.3.dev23/torchdeltaflow.egg-info/scm_version.json +8 -0
- torchdeltaflow-0.2.2/deltaflow/samplers/__init__.py +0 -7
- torchdeltaflow-0.2.2/torchdeltaflow.egg-info/scm_version.json +0 -8
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/.github/workflows/ci.yml +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/.github/workflows/docs.yml +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/.gitignore +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/CITATION.cff +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/CODE_OF_CONDUCT.md +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/CONTRIBUTING.md +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/LICENSE +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/MANIFEST.in +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/README.md +3 -3
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/benchmarks/README.md +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/core/base.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/core/base_interpolant.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/core/base_loss.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/core/base_solver.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/core/base_velocity_field.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/datasets/__init__.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/datasets/radiograph.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/interpolants/base.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/interpolants/schrodinger_bridge.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/interpolants/variance_preserving.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/inverse/__init__.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/inverse/likelihood.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/inverse/tweedie.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/losses/conditional_flow_matching.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/losses/delta_alignment.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/losses/flow_matching.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/models/backbone.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/models/ema.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/models/projector.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/samplers/euler.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/solvers/euler.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/solvers/heun.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/solvers/posterior_solver.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/trainer/__init__.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/trainer/data.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/trainer/loop.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/deltaflow/utils/numerical.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/api/core.md +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/api/datasets.md +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/api/index.md +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/api/interpolants.md +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/api/inverse.md +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/api/losses.md +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/api/models.md +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/api/samplers.md +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/api/solvers.md +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/api/trainer.md +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/api/utils.md +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/algorithm-comparison/comparison.gif +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/algorithm-comparison/comparison.png +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/algorithm-comparison/trajectories_comparison.png +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/css/extra.css +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/favicon.ico +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/favicon.png +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/guidance-alignment/features.png +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/inverse-posterior/inverse_posterior.gif +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/inverse-posterior/inverse_posterior.png +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/js/mathjax.js +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/landmark-detection/landmark_detection.png +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/logo.svg +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/minibatch-ot/minibatch_ot.gif +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/minibatch-ot/minibatch_ot.png +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/sampling-flow/flow.gif +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/sampling-flow/snapshots.png +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/sampling-flow/trajectories.png +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/sampling-flow/velocity_field.png +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/schrodinger-bridge/bridge_paths.gif +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/schrodinger-bridge/bridge_paths.png +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/schrodinger-bridge/sb_flow.gif +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/schrodinger-bridge/sb_snapshots.png +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/assets/schrodinger-bridge/sb_trajectories.png +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/concepts/architecture.md +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/contributing.md +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/getting-started/quickstart.md +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/guides/delta-alignment.md +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/guides/flow-matching.md +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/guides/inverse-problems.md +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/docs/index.md +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/examples/00-foundations/01-linear-interpolant/main.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/examples/10-sampling/01-euler-flow/main.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/examples/20-training/01-flow-matching/main.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/examples/20-training/02-delta-alignment/main.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/examples/30-inverse/01-posterior/main.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/examples/90-showcase/01-landmark-viz/main.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/examples/90-showcase/02-sampling-flow-viz/main.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/examples/90-showcase/04-inverse-posterior-viz/main.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/examples/90-showcase/05-schrodinger-bridge-viz/main.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/examples/90-showcase/06-algorithm-comparison/main.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/examples/90-showcase/07-landmark-detection/main.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/examples/90-showcase/08-guidance-alignment-pretraining/main.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/examples/90-showcase/README.md +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/mkdocs.yml +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/overrides/main.html +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/requirements-docs.txt +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/requirements.txt +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/setup.cfg +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/tests/__init__.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/tests/conftest.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/tests/test_interpolants.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/tests/test_interpolants_new.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/tests/test_inverse.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/tests/test_losses.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/tests/test_models.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/tests/test_samplers.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/tests/test_solvers.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/tests/test_trainer.py +0 -0
- {torchdeltaflow-0.2.2 → torchdeltaflow-0.2.3.dev23}/torchdeltaflow.egg-info/dependency_links.txt +0 -0
- {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
|
|
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
|
-
#
|
|
8
|
-
#
|
|
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.
|
|
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,
|
|
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
|
|
14
|
-
|
|
15
|
-
|
|
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
|
-
|
|
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
|
|
95
|
-
``scipy`` is
|
|
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 =
|
|
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
|
-
|
|
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
|
]
|