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.
- torchdeltaflow-0.1.0/.github/workflows/ci.yml +25 -0
- torchdeltaflow-0.1.0/.gitignore +79 -0
- torchdeltaflow-0.1.0/CHANGELOG.md +16 -0
- torchdeltaflow-0.1.0/CITATION.cff +11 -0
- torchdeltaflow-0.1.0/CODE_OF_CONDUCT.md +28 -0
- torchdeltaflow-0.1.0/CONTRIBUTING.md +26 -0
- torchdeltaflow-0.1.0/LICENSE +21 -0
- torchdeltaflow-0.1.0/MANIFEST.in +4 -0
- torchdeltaflow-0.1.0/PKG-INFO +169 -0
- torchdeltaflow-0.1.0/README.md +93 -0
- torchdeltaflow-0.1.0/benchmarks/README.md +7 -0
- torchdeltaflow-0.1.0/deltaflow/__init__.py +40 -0
- torchdeltaflow-0.1.0/deltaflow/core/__init__.py +18 -0
- torchdeltaflow-0.1.0/deltaflow/core/base.py +8 -0
- torchdeltaflow-0.1.0/deltaflow/core/base_interpolant.py +26 -0
- torchdeltaflow-0.1.0/deltaflow/core/base_loss.py +19 -0
- torchdeltaflow-0.1.0/deltaflow/core/base_solver.py +76 -0
- torchdeltaflow-0.1.0/deltaflow/core/base_velocity_field.py +20 -0
- torchdeltaflow-0.1.0/deltaflow/datasets/__init__.py +10 -0
- torchdeltaflow-0.1.0/deltaflow/datasets/radiograph.py +97 -0
- torchdeltaflow-0.1.0/deltaflow/interpolants/__init__.py +13 -0
- torchdeltaflow-0.1.0/deltaflow/interpolants/base.py +7 -0
- torchdeltaflow-0.1.0/deltaflow/interpolants/linear.py +32 -0
- torchdeltaflow-0.1.0/deltaflow/interpolants/ot.py +89 -0
- torchdeltaflow-0.1.0/deltaflow/interpolants/variance_preserving.py +40 -0
- torchdeltaflow-0.1.0/deltaflow/inverse/__init__.py +35 -0
- torchdeltaflow-0.1.0/deltaflow/inverse/likelihood.py +73 -0
- torchdeltaflow-0.1.0/deltaflow/inverse/operators.py +93 -0
- torchdeltaflow-0.1.0/deltaflow/inverse/tweedie.py +91 -0
- torchdeltaflow-0.1.0/deltaflow/losses/__init__.py +11 -0
- torchdeltaflow-0.1.0/deltaflow/losses/conditional_flow_matching.py +74 -0
- torchdeltaflow-0.1.0/deltaflow/losses/delta_alignment.py +139 -0
- torchdeltaflow-0.1.0/deltaflow/losses/flow_matching.py +7 -0
- torchdeltaflow-0.1.0/deltaflow/models/__init__.py +13 -0
- torchdeltaflow-0.1.0/deltaflow/models/backbone.py +110 -0
- torchdeltaflow-0.1.0/deltaflow/models/ema.py +28 -0
- torchdeltaflow-0.1.0/deltaflow/models/projector.py +96 -0
- torchdeltaflow-0.1.0/deltaflow/samplers/__init__.py +7 -0
- torchdeltaflow-0.1.0/deltaflow/samplers/euler.py +8 -0
- torchdeltaflow-0.1.0/deltaflow/solvers/__init__.py +11 -0
- torchdeltaflow-0.1.0/deltaflow/solvers/euler.py +61 -0
- torchdeltaflow-0.1.0/deltaflow/solvers/heun.py +34 -0
- torchdeltaflow-0.1.0/deltaflow/solvers/posterior_solver.py +102 -0
- torchdeltaflow-0.1.0/deltaflow/trainer/__init__.py +18 -0
- torchdeltaflow-0.1.0/deltaflow/trainer/coupling.py +57 -0
- torchdeltaflow-0.1.0/deltaflow/trainer/data.py +104 -0
- torchdeltaflow-0.1.0/deltaflow/trainer/loop.py +269 -0
- torchdeltaflow-0.1.0/deltaflow/utils/__init__.py +23 -0
- torchdeltaflow-0.1.0/deltaflow/utils/numerical.py +66 -0
- torchdeltaflow-0.1.0/docs/getting-started.md +42 -0
- torchdeltaflow-0.1.0/docs/index.md +15 -0
- torchdeltaflow-0.1.0/examples/00-foundations/01-linear-interpolant/main.py +34 -0
- torchdeltaflow-0.1.0/examples/10-sampling/01-euler-flow/main.py +60 -0
- torchdeltaflow-0.1.0/examples/20-training/01-flow-matching/main.py +47 -0
- torchdeltaflow-0.1.0/examples/20-training/02-delta-alignment/main.py +51 -0
- torchdeltaflow-0.1.0/examples/30-inverse/01-posterior/main.py +146 -0
- torchdeltaflow-0.1.0/examples/90-showcase/README.md +7 -0
- torchdeltaflow-0.1.0/mkdocs.yml +26 -0
- torchdeltaflow-0.1.0/pyproject.toml +110 -0
- torchdeltaflow-0.1.0/requirements-docs.txt +3 -0
- torchdeltaflow-0.1.0/requirements.txt +4 -0
- torchdeltaflow-0.1.0/setup.cfg +4 -0
- torchdeltaflow-0.1.0/tests/__init__.py +0 -0
- torchdeltaflow-0.1.0/tests/conftest.py +18 -0
- torchdeltaflow-0.1.0/tests/test_coupling.py +58 -0
- torchdeltaflow-0.1.0/tests/test_interpolants.py +28 -0
- torchdeltaflow-0.1.0/tests/test_interpolants_new.py +62 -0
- torchdeltaflow-0.1.0/tests/test_inverse.py +94 -0
- torchdeltaflow-0.1.0/tests/test_losses.py +31 -0
- torchdeltaflow-0.1.0/tests/test_models.py +38 -0
- torchdeltaflow-0.1.0/tests/test_samplers.py +26 -0
- torchdeltaflow-0.1.0/tests/test_solvers.py +91 -0
- torchdeltaflow-0.1.0/tests/test_trainer.py +65 -0
- torchdeltaflow-0.1.0/torchdeltaflow.egg-info/PKG-INFO +169 -0
- torchdeltaflow-0.1.0/torchdeltaflow.egg-info/SOURCES.txt +78 -0
- torchdeltaflow-0.1.0/torchdeltaflow.egg-info/dependency_links.txt +1 -0
- torchdeltaflow-0.1.0/torchdeltaflow.egg-info/requires.txt +28 -0
- torchdeltaflow-0.1.0/torchdeltaflow.egg-info/scm_file_list.json +75 -0
- torchdeltaflow-0.1.0/torchdeltaflow.egg-info/scm_version.json +8 -0
- 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,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,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
|