torchdeltaflow 0.1.0__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 (80) hide show
  1. torchdeltaflow-0.1.0/.github/workflows/ci.yml +25 -0
  2. torchdeltaflow-0.1.0/.gitignore +79 -0
  3. torchdeltaflow-0.1.0/CHANGELOG.md +16 -0
  4. torchdeltaflow-0.1.0/CITATION.cff +11 -0
  5. torchdeltaflow-0.1.0/CODE_OF_CONDUCT.md +28 -0
  6. torchdeltaflow-0.1.0/CONTRIBUTING.md +26 -0
  7. torchdeltaflow-0.1.0/LICENSE +21 -0
  8. torchdeltaflow-0.1.0/MANIFEST.in +4 -0
  9. torchdeltaflow-0.1.0/PKG-INFO +169 -0
  10. torchdeltaflow-0.1.0/README.md +93 -0
  11. torchdeltaflow-0.1.0/benchmarks/README.md +7 -0
  12. torchdeltaflow-0.1.0/deltaflow/__init__.py +40 -0
  13. torchdeltaflow-0.1.0/deltaflow/core/__init__.py +18 -0
  14. torchdeltaflow-0.1.0/deltaflow/core/base.py +8 -0
  15. torchdeltaflow-0.1.0/deltaflow/core/base_interpolant.py +26 -0
  16. torchdeltaflow-0.1.0/deltaflow/core/base_loss.py +19 -0
  17. torchdeltaflow-0.1.0/deltaflow/core/base_solver.py +76 -0
  18. torchdeltaflow-0.1.0/deltaflow/core/base_velocity_field.py +20 -0
  19. torchdeltaflow-0.1.0/deltaflow/datasets/__init__.py +10 -0
  20. torchdeltaflow-0.1.0/deltaflow/datasets/radiograph.py +97 -0
  21. torchdeltaflow-0.1.0/deltaflow/interpolants/__init__.py +13 -0
  22. torchdeltaflow-0.1.0/deltaflow/interpolants/base.py +7 -0
  23. torchdeltaflow-0.1.0/deltaflow/interpolants/linear.py +32 -0
  24. torchdeltaflow-0.1.0/deltaflow/interpolants/ot.py +89 -0
  25. torchdeltaflow-0.1.0/deltaflow/interpolants/variance_preserving.py +40 -0
  26. torchdeltaflow-0.1.0/deltaflow/inverse/__init__.py +35 -0
  27. torchdeltaflow-0.1.0/deltaflow/inverse/likelihood.py +73 -0
  28. torchdeltaflow-0.1.0/deltaflow/inverse/operators.py +93 -0
  29. torchdeltaflow-0.1.0/deltaflow/inverse/tweedie.py +91 -0
  30. torchdeltaflow-0.1.0/deltaflow/losses/__init__.py +11 -0
  31. torchdeltaflow-0.1.0/deltaflow/losses/conditional_flow_matching.py +74 -0
  32. torchdeltaflow-0.1.0/deltaflow/losses/delta_alignment.py +139 -0
  33. torchdeltaflow-0.1.0/deltaflow/losses/flow_matching.py +7 -0
  34. torchdeltaflow-0.1.0/deltaflow/models/__init__.py +13 -0
  35. torchdeltaflow-0.1.0/deltaflow/models/backbone.py +110 -0
  36. torchdeltaflow-0.1.0/deltaflow/models/ema.py +28 -0
  37. torchdeltaflow-0.1.0/deltaflow/models/projector.py +96 -0
  38. torchdeltaflow-0.1.0/deltaflow/samplers/__init__.py +7 -0
  39. torchdeltaflow-0.1.0/deltaflow/samplers/euler.py +8 -0
  40. torchdeltaflow-0.1.0/deltaflow/solvers/__init__.py +11 -0
  41. torchdeltaflow-0.1.0/deltaflow/solvers/euler.py +61 -0
  42. torchdeltaflow-0.1.0/deltaflow/solvers/heun.py +34 -0
  43. torchdeltaflow-0.1.0/deltaflow/solvers/posterior_solver.py +102 -0
  44. torchdeltaflow-0.1.0/deltaflow/trainer/__init__.py +18 -0
  45. torchdeltaflow-0.1.0/deltaflow/trainer/coupling.py +57 -0
  46. torchdeltaflow-0.1.0/deltaflow/trainer/data.py +104 -0
  47. torchdeltaflow-0.1.0/deltaflow/trainer/loop.py +269 -0
  48. torchdeltaflow-0.1.0/deltaflow/utils/__init__.py +23 -0
  49. torchdeltaflow-0.1.0/deltaflow/utils/numerical.py +66 -0
  50. torchdeltaflow-0.1.0/docs/getting-started.md +42 -0
  51. torchdeltaflow-0.1.0/docs/index.md +15 -0
  52. torchdeltaflow-0.1.0/examples/00-foundations/01-linear-interpolant/main.py +34 -0
  53. torchdeltaflow-0.1.0/examples/10-sampling/01-euler-flow/main.py +60 -0
  54. torchdeltaflow-0.1.0/examples/20-training/01-flow-matching/main.py +47 -0
  55. torchdeltaflow-0.1.0/examples/20-training/02-delta-alignment/main.py +51 -0
  56. torchdeltaflow-0.1.0/examples/30-inverse/01-posterior/main.py +146 -0
  57. torchdeltaflow-0.1.0/examples/90-showcase/README.md +7 -0
  58. torchdeltaflow-0.1.0/mkdocs.yml +26 -0
  59. torchdeltaflow-0.1.0/pyproject.toml +110 -0
  60. torchdeltaflow-0.1.0/requirements-docs.txt +3 -0
  61. torchdeltaflow-0.1.0/requirements.txt +4 -0
  62. torchdeltaflow-0.1.0/setup.cfg +4 -0
  63. torchdeltaflow-0.1.0/tests/__init__.py +0 -0
  64. torchdeltaflow-0.1.0/tests/conftest.py +18 -0
  65. torchdeltaflow-0.1.0/tests/test_coupling.py +58 -0
  66. torchdeltaflow-0.1.0/tests/test_interpolants.py +28 -0
  67. torchdeltaflow-0.1.0/tests/test_interpolants_new.py +62 -0
  68. torchdeltaflow-0.1.0/tests/test_inverse.py +94 -0
  69. torchdeltaflow-0.1.0/tests/test_losses.py +31 -0
  70. torchdeltaflow-0.1.0/tests/test_models.py +38 -0
  71. torchdeltaflow-0.1.0/tests/test_samplers.py +26 -0
  72. torchdeltaflow-0.1.0/tests/test_solvers.py +91 -0
  73. torchdeltaflow-0.1.0/tests/test_trainer.py +65 -0
  74. torchdeltaflow-0.1.0/torchdeltaflow.egg-info/PKG-INFO +169 -0
  75. torchdeltaflow-0.1.0/torchdeltaflow.egg-info/SOURCES.txt +78 -0
  76. torchdeltaflow-0.1.0/torchdeltaflow.egg-info/dependency_links.txt +1 -0
  77. torchdeltaflow-0.1.0/torchdeltaflow.egg-info/requires.txt +28 -0
  78. torchdeltaflow-0.1.0/torchdeltaflow.egg-info/scm_file_list.json +75 -0
  79. torchdeltaflow-0.1.0/torchdeltaflow.egg-info/scm_version.json +8 -0
  80. torchdeltaflow-0.1.0/torchdeltaflow.egg-info/top_level.txt +1 -0
@@ -0,0 +1,25 @@
1
+ name: CI
2
+
3
+ on:
4
+ push:
5
+ branches: [main]
6
+ pull_request:
7
+ branches: [main]
8
+
9
+ jobs:
10
+ test:
11
+ runs-on: ubuntu-latest
12
+ strategy:
13
+ matrix:
14
+ python-version: ["3.10", "3.11", "3.12"]
15
+ steps:
16
+ - uses: actions/checkout@v4
17
+ - uses: actions/setup-python@v5
18
+ with:
19
+ python-version: ${{ matrix.python-version }}
20
+ - name: Install
21
+ run: |
22
+ python -m pip install --upgrade pip
23
+ pip install -e ".[dev]"
24
+ - name: Test
25
+ run: pytest --cov=deltaflow
@@ -0,0 +1,79 @@
1
+ # --- Python ---
2
+ __pycache__/
3
+ *.py[cod]
4
+ *.pyo
5
+ *.pyd
6
+ *$py.class
7
+ .Python
8
+
9
+ # --- Build / packaging ---
10
+ build/
11
+ dist/
12
+ wheels/
13
+ *.egg-info/
14
+ *.egg
15
+ .eggs/
16
+ pip-wheel-metadata/
17
+ MANIFEST
18
+
19
+ # setuptools_scm writes deltaflow/_version.py; keep it out of git
20
+ deltaflow/_version.py
21
+
22
+ # --- Virtual environments ---
23
+ venv/
24
+ .venv/
25
+ env/
26
+ .env
27
+ .env.*
28
+
29
+ # --- Testing / coverage ---
30
+ .pytest_cache/
31
+ .coverage
32
+ .coverage.*
33
+ htmlcov/
34
+ coverage.xml
35
+ .tox/
36
+ .nox/
37
+ .hypothesis/
38
+
39
+ # --- Type checkers / linters ---
40
+ .mypy_cache/
41
+ .pyright/
42
+ .ruff_cache/
43
+ .pytype/
44
+
45
+ # --- Docs ---
46
+ site/
47
+ docs/_build/
48
+
49
+ # --- Jupyter ---
50
+ .ipynb_checkpoints/
51
+
52
+ # --- Training artifacts (created by deltaflow.trainer + examples) ---
53
+ checkpoints/
54
+ outputs/
55
+ runs/
56
+ wandb/
57
+ logs/
58
+ *.pt
59
+ *.pth
60
+ *.ckpt
61
+ *.safetensors
62
+
63
+ # --- Data ---
64
+ data/
65
+ datasets/local/
66
+ *.h5
67
+ *.hdf5
68
+ *.npy
69
+ *.npz
70
+
71
+ # --- OS / editor ---
72
+ .DS_Store
73
+ Thumbs.db
74
+ Desktop.ini
75
+ *.swp
76
+ *.swo
77
+ .idea/
78
+ .vscode/
79
+ *.log
@@ -0,0 +1,16 @@
1
+ # Changelog
2
+
3
+ All notable changes to this project are documented here.
4
+ The format follows [Keep a Changelog](https://keepachangelog.com/en/1.1.0/).
5
+
6
+ ## [Unreleased]
7
+
8
+ ### Added
9
+
10
+ - Initial scaffold: `core`, `interpolants` (`LinearInterpolant`), `samplers`
11
+ (`FlowSampler`), `losses` (`FlowMatchingLoss`, `DeltaAlignmentLoss`),
12
+ `models` (`MultiScaleProjector`, `EMA`), `datasets` (generic radiograph
13
+ wrappers).
14
+ - Tiered examples (`00-foundations`, `10-sampling`, `20-training`,
15
+ `90-showcase`).
16
+ - Unit tests for interpolants, samplers, losses, and models.
@@ -0,0 +1,11 @@
1
+ cff-version: 1.2.0
2
+ message: "If you use this software, please cite it as below."
3
+ title: "torchDeltaFlow: Flow Matching, Optimal-Transport Coupling, and Posterior Sampling for Inverse Problems in PyTorch"
4
+ authors:
5
+ - family-names: "Limbunlom"
6
+ given-names: "Phrugsa"
7
+ email: "phrugsa.lim@gmail.com"
8
+ date-released: 2026-07-22
9
+ url: "https://github.com/phrugsa-limbunlom/deltaflow"
10
+ repository-code: "https://github.com/phrugsa-limbunlom/deltaflow"
11
+ license: MIT
@@ -0,0 +1,28 @@
1
+ # Contributor Covenant Code of Conduct
2
+
3
+ ## Our Pledge
4
+
5
+ We as members, contributors, and leaders pledge to make participation in our
6
+ community a harassment-free experience for everyone, regardless of age, body
7
+ size, visible or invisible disability, ethnicity, sex characteristics,
8
+ gender identity and expression, level of experience, education,
9
+ socio-economic status, nationality, personal appearance, race, religion, or
10
+ sexual identity and orientation.
11
+
12
+ ## Our Standards
13
+
14
+ Examples of behavior that contributes to a positive environment include
15
+ being respectful of differing opinions, giving and gracefully accepting
16
+ constructive feedback, and focusing on what is best for the community.
17
+
18
+ Examples of unacceptable behavior include harassment, insulting or
19
+ derogatory comments, and other conduct which could reasonably be considered
20
+ inappropriate in a professional setting.
21
+
22
+ ## Enforcement
23
+
24
+ Instances of abusive, harassing, or otherwise unacceptable behavior may be
25
+ reported to the project maintainers. All complaints will be reviewed and
26
+ investigated promptly and fairly.
27
+
28
+ This Code of Conduct is adapted from the [Contributor Covenant](https://www.contributor-covenant.org/), version 2.1.
@@ -0,0 +1,26 @@
1
+ # Contributing
2
+
3
+ Contributions are welcome. This project is small and research-oriented, so
4
+ please open an issue to discuss non-trivial changes before sending a PR.
5
+
6
+ ## Setup
7
+
8
+ ```bash
9
+ git clone https://github.com/yourname/deltaflow
10
+ cd deltaflow
11
+ pip install -e ".[dev]"
12
+ ```
13
+
14
+ ## Before opening a PR
15
+
16
+ ```bash
17
+ black deltaflow tests examples
18
+ isort deltaflow tests examples
19
+ pytest
20
+ ```
21
+
22
+ ## Adding an example
23
+
24
+ Examples live under `examples/<tier>/<NN-name>/main.py`, where `<tier>` is
25
+ one of `00-foundations`, `10-sampling`, `20-training`, `90-showcase`. Each
26
+ example should run standalone in well under a minute on CPU.
@@ -0,0 +1,21 @@
1
+ MIT License
2
+
3
+ Copyright (c) 2026 DeltaFlow 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,4 @@
1
+ include README.md
2
+ include LICENSE
3
+ include CHANGELOG.md
4
+ recursive-include deltaflow *.py
@@ -0,0 +1,169 @@
1
+ Metadata-Version: 2.4
2
+ Name: torchdeltaflow
3
+ Version: 0.1.0
4
+ Summary: Flow matching, mini-batch optimal transport coupling, and posterior sampling for inverse problems in PyTorch
5
+ Author-email: Phrugsa Limbunlom <phrugsa.lim@gmail.com>
6
+ Maintainer-email: Phrugsa Limbunlom <phrugsa.lim@gmail.com>
7
+ License: MIT License
8
+
9
+ Copyright (c) 2026 DeltaFlow Contributors
10
+
11
+ Permission is hereby granted, free of charge, to any person obtaining a copy
12
+ of this software and associated documentation files (the "Software"), to deal
13
+ in the Software without restriction, including without limitation the rights
14
+ to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
15
+ copies of the Software, and to permit persons to whom the Software is
16
+ furnished to do so, subject to the following conditions:
17
+
18
+ The above copyright notice and this permission notice shall be included in all
19
+ copies or substantial portions of the Software.
20
+
21
+ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
22
+ IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
23
+ FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
24
+ AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
25
+ LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
26
+ OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
27
+ SOFTWARE.
28
+
29
+ Project-URL: Homepage, https://github.com/phrugsa-limbunlom/deltaflow
30
+ Project-URL: Repository, https://github.com/phrugsa-limbunlom/deltaflow
31
+ Project-URL: Issues, https://github.com/phrugsa-limbunlom/deltaflow/issues
32
+ Project-URL: Changelog, https://github.com/phrugsa-limbunlom/deltaflow/blob/main/CHANGELOG.md
33
+ Keywords: pytorch,flow-matching,rectified-flow,optimal-transport,diffusion-models,generative-models,posterior-sampling,inverse-problems,medical-imaging,deep-learning
34
+ Classifier: Development Status :: 3 - Alpha
35
+ Classifier: Intended Audience :: Developers
36
+ Classifier: Intended Audience :: Science/Research
37
+ Classifier: License :: OSI Approved :: MIT License
38
+ Classifier: Operating System :: OS Independent
39
+ Classifier: Programming Language :: Python :: 3
40
+ Classifier: Programming Language :: Python :: 3.9
41
+ Classifier: Programming Language :: Python :: 3.10
42
+ Classifier: Programming Language :: Python :: 3.11
43
+ Classifier: Programming Language :: Python :: 3.12
44
+ Classifier: Programming Language :: Python :: 3.13
45
+ Classifier: Topic :: Scientific/Engineering :: Artificial Intelligence
46
+ Classifier: Topic :: Scientific/Engineering :: Medical Science Apps.
47
+ Classifier: Topic :: Software Development :: Libraries
48
+ Classifier: Topic :: Software Development :: Libraries :: Python Modules
49
+ Requires-Python: >=3.9
50
+ Description-Content-Type: text/markdown
51
+ License-File: LICENSE
52
+ Requires-Dist: torch>=2.0
53
+ Requires-Dist: numpy
54
+ Requires-Dist: einops
55
+ Requires-Dist: tqdm
56
+ Provides-Extra: images
57
+ Requires-Dist: pillow>=9.0; extra == "images"
58
+ Provides-Extra: ot
59
+ Requires-Dist: scipy>=1.10; extra == "ot"
60
+ Provides-Extra: all
61
+ Requires-Dist: pillow>=9.0; extra == "all"
62
+ Requires-Dist: scipy>=1.10; extra == "all"
63
+ Provides-Extra: dev
64
+ Requires-Dist: pytest>=6.0; extra == "dev"
65
+ Requires-Dist: pytest-cov>=2.0; extra == "dev"
66
+ Requires-Dist: black>=24.0; extra == "dev"
67
+ Requires-Dist: isort>=5.0; extra == "dev"
68
+ Requires-Dist: mypy>=1.0; extra == "dev"
69
+ Requires-Dist: pillow>=9.0; extra == "dev"
70
+ Requires-Dist: scipy>=1.10; extra == "dev"
71
+ Provides-Extra: docs
72
+ Requires-Dist: mkdocs-material>=9.0.0; extra == "docs"
73
+ Requires-Dist: mkdocstrings[python]>=0.18; extra == "docs"
74
+ Requires-Dist: pymdown-extensions>=9.0; extra == "docs"
75
+ Dynamic: license-file
76
+
77
+ # DeltaFlow
78
+
79
+ **Flow matching and anatomy-invariant guidance alignment for radiograph generative pretraining.**
80
+
81
+ `DeltaFlow` provides composable PyTorch primitives for two things that are
82
+ usually bundled together in ad-hoc research code:
83
+
84
+ - **Flow matching** — a simulation-free generative objective that regresses
85
+ a velocity field onto the conditional velocity of a probability path
86
+ (Lipman et al., 2023), sampled with a simple Euler ODE solver.
87
+ - **Delta alignment** — a multi-scale, anatomy-cancelling loss that aligns
88
+ the *guidance-difference* feature `Δh = h_cond - h_uncond` across
89
+ augmented views, instead of aligning raw (anatomy-entangled) features.
90
+ This is the mechanism behind the name: **Delta** is the guidance-difference
91
+ feature, **Flow** is the flow-matching engine that consumes it.
92
+
93
+ The library targets 2D radiography (chest X-ray, cephalometric, and hand
94
+ radiographs), but the core primitives (`interpolants`, `samplers`, `losses`,
95
+ `models.projector`) are domain-agnostic.
96
+
97
+ ## What's inside
98
+
99
+ | Module | Role |
100
+ |---|---|
101
+ | `deltaflow.core` | `BaseVelocityField` — the interface every backbone implements |
102
+ | `deltaflow.interpolants` | Probability paths, e.g. `LinearInterpolant` (rectified flow) |
103
+ | `deltaflow.samplers` | `FlowSampler` — Euler ODE integration |
104
+ | `deltaflow.losses` | `FlowMatchingLoss`, `DeltaAlignmentLoss` |
105
+ | `deltaflow.models` | `MultiScaleProjector` (per-level projection heads), `EMA` |
106
+ | `deltaflow.datasets` | Generic radiograph dataset wrappers (chest / cephalo / hand) |
107
+
108
+ ## Install
109
+
110
+ ```bash
111
+ pip install -e .
112
+ # or, for docs/dev tooling:
113
+ pip install -e ".[dev]"
114
+ ```
115
+
116
+ ## Usage
117
+
118
+ ### Flow matching
119
+
120
+ ```python
121
+ import torch
122
+ from deltaflow.interpolants import LinearInterpolant
123
+ from deltaflow.losses import FlowMatchingLoss
124
+ from deltaflow.samplers import FlowSampler
125
+
126
+ loss_fn = FlowMatchingLoss(interpolant=LinearInterpolant())
127
+ loss = loss_fn(model, x1) # model(x_t, t) -> predicted velocity
128
+ loss.backward()
129
+
130
+ samples = FlowSampler(model).sample(torch.randn(1000, 2), n_steps=50)
131
+ ```
132
+
133
+ ### Delta alignment
134
+
135
+ ```python
136
+ from deltaflow.losses.delta_alignment import DeltaAlignmentLoss
137
+ from deltaflow.models import MultiScaleProjector
138
+
139
+ projector = MultiScaleProjector(feature_dims={"enc_1_4": 256, "bottleneck": 1024})
140
+ loss_fn = DeltaAlignmentLoss(projector, lambda_flow=1.0, lambda_align=5.0)
141
+
142
+ total, loss_dict = loss_fn(
143
+ v_c1, v_u1, target_v1,
144
+ v_c2, v_u2, target_v2,
145
+ feats_u1, feats_c1, feats_u2, feats_c2,
146
+ )
147
+ ```
148
+
149
+ See [`examples/`](examples/) for runnable, tiered walkthroughs
150
+ (`00-foundations/`, `10-sampling/`, `20-training/`, `90-showcase/`).
151
+
152
+ ## Development
153
+
154
+ ```bash
155
+ pip install -e ".[dev]"
156
+ pytest
157
+ ```
158
+
159
+ ## Contributing
160
+
161
+ See [CONTRIBUTING.md](CONTRIBUTING.md).
162
+
163
+ ## License
164
+
165
+ MIT. See [LICENSE](LICENSE).
166
+
167
+ ## Citation
168
+
169
+ If you use DeltaFlow in your research, please cite it (see [CITATION.cff](CITATION.cff)).
@@ -0,0 +1,93 @@
1
+ # DeltaFlow
2
+
3
+ **Flow matching and anatomy-invariant guidance alignment for radiograph generative pretraining.**
4
+
5
+ `DeltaFlow` provides composable PyTorch primitives for two things that are
6
+ usually bundled together in ad-hoc research code:
7
+
8
+ - **Flow matching** — a simulation-free generative objective that regresses
9
+ a velocity field onto the conditional velocity of a probability path
10
+ (Lipman et al., 2023), sampled with a simple Euler ODE solver.
11
+ - **Delta alignment** — a multi-scale, anatomy-cancelling loss that aligns
12
+ the *guidance-difference* feature `Δh = h_cond - h_uncond` across
13
+ augmented views, instead of aligning raw (anatomy-entangled) features.
14
+ This is the mechanism behind the name: **Delta** is the guidance-difference
15
+ feature, **Flow** is the flow-matching engine that consumes it.
16
+
17
+ The library targets 2D radiography (chest X-ray, cephalometric, and hand
18
+ radiographs), but the core primitives (`interpolants`, `samplers`, `losses`,
19
+ `models.projector`) are domain-agnostic.
20
+
21
+ ## What's inside
22
+
23
+ | Module | Role |
24
+ |---|---|
25
+ | `deltaflow.core` | `BaseVelocityField` — the interface every backbone implements |
26
+ | `deltaflow.interpolants` | Probability paths, e.g. `LinearInterpolant` (rectified flow) |
27
+ | `deltaflow.samplers` | `FlowSampler` — Euler ODE integration |
28
+ | `deltaflow.losses` | `FlowMatchingLoss`, `DeltaAlignmentLoss` |
29
+ | `deltaflow.models` | `MultiScaleProjector` (per-level projection heads), `EMA` |
30
+ | `deltaflow.datasets` | Generic radiograph dataset wrappers (chest / cephalo / hand) |
31
+
32
+ ## Install
33
+
34
+ ```bash
35
+ pip install -e .
36
+ # or, for docs/dev tooling:
37
+ pip install -e ".[dev]"
38
+ ```
39
+
40
+ ## Usage
41
+
42
+ ### Flow matching
43
+
44
+ ```python
45
+ import torch
46
+ from deltaflow.interpolants import LinearInterpolant
47
+ from deltaflow.losses import FlowMatchingLoss
48
+ from deltaflow.samplers import FlowSampler
49
+
50
+ loss_fn = FlowMatchingLoss(interpolant=LinearInterpolant())
51
+ loss = loss_fn(model, x1) # model(x_t, t) -> predicted velocity
52
+ loss.backward()
53
+
54
+ samples = FlowSampler(model).sample(torch.randn(1000, 2), n_steps=50)
55
+ ```
56
+
57
+ ### Delta alignment
58
+
59
+ ```python
60
+ from deltaflow.losses.delta_alignment import DeltaAlignmentLoss
61
+ from deltaflow.models import MultiScaleProjector
62
+
63
+ projector = MultiScaleProjector(feature_dims={"enc_1_4": 256, "bottleneck": 1024})
64
+ loss_fn = DeltaAlignmentLoss(projector, lambda_flow=1.0, lambda_align=5.0)
65
+
66
+ total, loss_dict = loss_fn(
67
+ v_c1, v_u1, target_v1,
68
+ v_c2, v_u2, target_v2,
69
+ feats_u1, feats_c1, feats_u2, feats_c2,
70
+ )
71
+ ```
72
+
73
+ See [`examples/`](examples/) for runnable, tiered walkthroughs
74
+ (`00-foundations/`, `10-sampling/`, `20-training/`, `90-showcase/`).
75
+
76
+ ## Development
77
+
78
+ ```bash
79
+ pip install -e ".[dev]"
80
+ pytest
81
+ ```
82
+
83
+ ## Contributing
84
+
85
+ See [CONTRIBUTING.md](CONTRIBUTING.md).
86
+
87
+ ## License
88
+
89
+ MIT. See [LICENSE](LICENSE).
90
+
91
+ ## Citation
92
+
93
+ If you use DeltaFlow in your research, please cite it (see [CITATION.cff](CITATION.cff)).
@@ -0,0 +1,7 @@
1
+ # Benchmarks
2
+
3
+ Placeholder for reproducible benchmark scripts (e.g. wall-clock/memory of
4
+ `FlowSampler` at different step counts, or convergence speed of
5
+ `DeltaAlignmentLoss` vs. a raw-feature alignment baseline). Follow
6
+ `torchebm`'s pattern: one script per benchmark, results written to a
7
+ tracked JSON/CSV so regressions are visible in CI.
@@ -0,0 +1,40 @@
1
+ """
2
+ DeltaFlow: a PyTorch library for flow matching, mini-batch optimal-transport
3
+ coupling, and posterior sampling for inverse problems.
4
+
5
+ Module map:
6
+
7
+ - :mod:`deltaflow.core` -- abstract base classes every component
8
+ subclasses (:class:`BaseVelocityField`, :class:`BaseInterpolant`,
9
+ :class:`BaseSolver`, :class:`BaseLoss`).
10
+ - :mod:`deltaflow.interpolants` -- probability paths (linear, mini-batch OT,
11
+ variance-preserving).
12
+ - :mod:`deltaflow.losses` -- conditional flow matching, plus the
13
+ optional delta-alignment loss for guidance-representation pretraining.
14
+ - :mod:`deltaflow.solvers` -- Euler, Heun, and the
15
+ :class:`~deltaflow.solvers.PosteriorSolver` which *wraps* a base solver
16
+ and injects the measurement-likelihood gradient per step (FlowDPS /
17
+ Flower style).
18
+ - :mod:`deltaflow.trainer` -- streaming image dataset, mixed-precision
19
+ training loop with grad accumulation and checkpoint/resume, and
20
+ train-time coupling strategies (independent vs OT).
21
+ - :mod:`deltaflow.inverse` -- measurement operators, Tweedie
22
+ decomposition, and Gaussian likelihood. v1 targets pixel space with an
23
+ optional decoder pullback for the latent-space case (see
24
+ ``deltaflow.inverse.__init__`` docstring).
25
+ - :mod:`deltaflow.models` -- backbone wrappers, EMA, and projector
26
+ heads used by the delta-alignment loss.
27
+ """
28
+
29
+ from importlib.metadata import PackageNotFoundError, version
30
+
31
+ try:
32
+ # The distribution name on PyPI is ``torchdeltaflow`` (the plain
33
+ # ``deltaflow`` name is claimed by an unrelated 2020 package). The
34
+ # *import* name is still ``deltaflow``, which is why we can't just
35
+ # call ``version(__name__)``.
36
+ __version__ = version("torchdeltaflow")
37
+ except PackageNotFoundError: # pragma: no cover - not installed
38
+ __version__ = "0.0.0.dev0"
39
+
40
+ __all__ = ["__version__"]
@@ -0,0 +1,18 @@
1
+ """Base abstractions shared across DeltaFlow's models, losses, and solvers.
2
+
3
+ Every user-facing component (velocity field, interpolant, solver, loss)
4
+ subclasses one of these bases, so new variants are drop-in and not
5
+ rewrites of the surrounding machinery.
6
+ """
7
+
8
+ from .base_interpolant import BaseInterpolant
9
+ from .base_loss import BaseLoss
10
+ from .base_solver import BaseSolver
11
+ from .base_velocity_field import BaseVelocityField
12
+
13
+ __all__ = [
14
+ "BaseInterpolant",
15
+ "BaseLoss",
16
+ "BaseSolver",
17
+ "BaseVelocityField",
18
+ ]
@@ -0,0 +1,8 @@
1
+ """Backward-compatibility shim. Prefer importing from the split modules
2
+ (:mod:`deltaflow.core.base_velocity_field`, etc.) or from :mod:`deltaflow.core`
3
+ directly.
4
+ """
5
+
6
+ from .base_velocity_field import BaseVelocityField
7
+
8
+ __all__ = ["BaseVelocityField"]
@@ -0,0 +1,26 @@
1
+ """Base class for probability-path interpolants."""
2
+
3
+ from abc import ABC, abstractmethod
4
+ from typing import Optional, Tuple
5
+
6
+ import torch
7
+
8
+
9
+ class BaseInterpolant(ABC):
10
+ """Base class for a probability path between noise ``x0`` and data ``x1``.
11
+
12
+ An interpolant defines, for each ``t in [0, 1]``, an intermediate point
13
+ ``x_t`` and its conditional target velocity ``u_t`` such that regressing
14
+ a model onto ``u_t`` (in expectation over the path) yields the marginal
15
+ velocity field of the flow-matching ODE.
16
+
17
+ Convention used throughout DeltaFlow: ``t = 0`` corresponds to noise
18
+ (``x_t = x_0``) and ``t = 1`` corresponds to data (``x_t = x_1``).
19
+ """
20
+
21
+ @abstractmethod
22
+ def interpolate(
23
+ self, x1: torch.Tensor, t: torch.Tensor, x0: Optional[torch.Tensor] = None
24
+ ) -> Tuple[torch.Tensor, torch.Tensor]:
25
+ """Return ``(x_t, target_velocity)`` for the given data/time (and optional noise)."""
26
+ raise NotImplementedError
@@ -0,0 +1,19 @@
1
+ """Base class for DeltaFlow training objectives."""
2
+
3
+ from abc import ABC, abstractmethod
4
+
5
+ import torch
6
+
7
+
8
+ class BaseLoss(ABC):
9
+ """Base class for a callable training loss.
10
+
11
+ Subclasses must implement :meth:`__call__` and return a scalar tensor
12
+ with ``requires_grad=True`` (assuming the model has trainable
13
+ parameters). The signature is intentionally flexible - individual losses
14
+ define which positional and keyword arguments they consume.
15
+ """
16
+
17
+ @abstractmethod
18
+ def __call__(self, *args, **kwargs) -> torch.Tensor:
19
+ raise NotImplementedError
@@ -0,0 +1,76 @@
1
+ """Base class for ODE/SDE solvers that integrate a learned velocity field."""
2
+
3
+ from abc import ABC, abstractmethod
4
+ from typing import Callable, Optional
5
+
6
+ import torch
7
+
8
+
9
+ class BaseSolver(ABC):
10
+ """Base class for numerical integrators of ``dx/dt = v_theta(x, t)``.
11
+
12
+ A solver holds a reference to a velocity model and exposes two methods:
13
+
14
+ - :meth:`step` performs a single integration step from ``(x, t)`` to
15
+ ``(x', t + dt)``. Subclasses implement the actual stepping rule.
16
+ - :meth:`sample` drives :meth:`step` in a loop from ``t_start`` to
17
+ ``t_end`` and returns the final state.
18
+
19
+ Design note: :class:`~deltaflow.solvers.posterior_solver.PosteriorSolver`
20
+ wraps a :class:`BaseSolver` and hooks the likelihood gradient into every
21
+ call to :meth:`step`, so the base stepping logic is never duplicated.
22
+
23
+ Args:
24
+ model: callable ``model(x, t, **cond) -> velocity`` with the same
25
+ signature as :class:`~deltaflow.core.BaseVelocityField`.
26
+ time_scale: multiplies the continuous ``t in [0, 1]`` before it is
27
+ passed to the model; useful when the backbone was trained with a
28
+ different numeric time convention (e.g. diffusion timesteps).
29
+ """
30
+
31
+ def __init__(self, model: Callable, time_scale: float = 1.0):
32
+ self.model = model
33
+ self.time_scale = time_scale
34
+
35
+ def _time_tensor(self, x: torch.Tensor, t: float) -> torch.Tensor:
36
+ return torch.full((x.shape[0],), t, device=x.device, dtype=x.dtype)
37
+
38
+ def _eval_velocity(self, x: torch.Tensor, t: float, **cond) -> torch.Tensor:
39
+ t_tensor = self._time_tensor(x, t) * self.time_scale
40
+ return self.model(x, t_tensor, **cond)
41
+
42
+ @abstractmethod
43
+ def step(self, x: torch.Tensor, t: float, dt: float, **cond) -> torch.Tensor:
44
+ """Advance the state ``x`` from time ``t`` to ``t + dt``."""
45
+ raise NotImplementedError
46
+
47
+ def sample(
48
+ self,
49
+ x: torch.Tensor,
50
+ n_steps: int = 50,
51
+ t_start: float = 0.0,
52
+ t_end: float = 1.0,
53
+ show_progress: bool = True,
54
+ progress_desc: Optional[str] = None,
55
+ **cond,
56
+ ) -> torch.Tensor:
57
+ """Integrate from ``t_start`` to ``t_end`` in ``n_steps`` uniform steps."""
58
+ dt = (t_end - t_start) / n_steps
59
+ steps = range(n_steps)
60
+ if show_progress:
61
+ try:
62
+ from tqdm import tqdm
63
+
64
+ steps = tqdm(
65
+ steps,
66
+ desc=progress_desc or type(self).__name__,
67
+ total=n_steps,
68
+ leave=False,
69
+ )
70
+ except ImportError:
71
+ pass
72
+
73
+ for i in steps:
74
+ t_val = t_start + i * dt
75
+ x = self.step(x, t_val, dt, **cond)
76
+ return x