torchdeltaflow 0.2.3.dev21__tar.gz → 0.2.4__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.4/.github/workflows/publish.yml +175 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/CHANGELOG.md +10 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/PKG-INFO +2 -2
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/README.md +1 -1
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/models/__init__.py +9 -1
- torchdeltaflow-0.2.4/deltaflow/models/dit.py +438 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/examples.md +36 -0
- torchdeltaflow-0.2.4/examples/20-training/03-conditional-dit/main.py +72 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/pyproject.toml +5 -4
- torchdeltaflow-0.2.4/tests/test_dit.py +147 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/torchdeltaflow.egg-info/PKG-INFO +2 -2
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/torchdeltaflow.egg-info/SOURCES.txt +3 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/torchdeltaflow.egg-info/scm_file_list.json +3 -0
- torchdeltaflow-0.2.4/torchdeltaflow.egg-info/scm_version.json +8 -0
- torchdeltaflow-0.2.3.dev21/.github/workflows/publish.yml +0 -78
- torchdeltaflow-0.2.3.dev21/torchdeltaflow.egg-info/scm_version.json +0 -8
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/.github/workflows/ci.yml +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/.github/workflows/docs.yml +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/.gitignore +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/CITATION.cff +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/CODE_OF_CONDUCT.md +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/CONTRIBUTING.md +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/LICENSE +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/MANIFEST.in +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/benchmarks/README.md +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/__init__.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/core/__init__.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/core/base.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/core/base_coupling.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/core/base_equilibrium_field.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/core/base_equilibrium_interpolant.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/core/base_interpolant.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/core/base_loss.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/core/base_solver.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/core/base_velocity_field.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/datasets/__init__.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/datasets/radiograph.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/interpolants/__init__.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/interpolants/base.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/interpolants/equilibrium.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/interpolants/linear.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/interpolants/ot.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/interpolants/schrodinger_bridge.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/interpolants/variance_preserving.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/inverse/__init__.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/inverse/likelihood.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/inverse/operators.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/inverse/tweedie.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/losses/__init__.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/losses/conditional_flow_matching.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/losses/delta_alignment.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/losses/equilibrium_matching.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/losses/flow_matching.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/models/backbone.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/models/ema.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/models/projector.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/samplers/__init__.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/samplers/euler.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/solvers/__init__.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/solvers/euler.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/solvers/gradient_descent.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/solvers/heun.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/solvers/posterior_solver.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/trainer/__init__.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/trainer/coupling.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/trainer/data.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/trainer/loop.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/utils/__init__.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/utils/numerical.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/utils/ot.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/api/core.md +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/api/datasets.md +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/api/index.md +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/api/interpolants.md +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/api/inverse.md +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/api/losses.md +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/api/models.md +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/api/samplers.md +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/api/solvers.md +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/api/trainer.md +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/api/utils.md +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/algorithm-comparison/comparison.gif +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/algorithm-comparison/comparison.png +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/algorithm-comparison/trajectories_comparison.png +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/css/extra.css +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/equilibrium-matching/energy_landscape.png +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/equilibrium-matching/eqm_sampling.gif +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/equilibrium-matching/gd_snapshots.png +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/equilibrium-matching/gd_trajectories.png +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/favicon.ico +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/favicon.png +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/guidance-alignment/features.png +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/inverse-posterior/inverse_posterior.gif +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/inverse-posterior/inverse_posterior.png +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/js/mathjax.js +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/landmark-detection/landmark_detection.png +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/logo.svg +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/minibatch-ot/minibatch_ot.gif +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/minibatch-ot/minibatch_ot.png +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/sampling-flow/flow.gif +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/sampling-flow/snapshots.png +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/sampling-flow/trajectories.png +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/sampling-flow/velocity_field.png +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/schrodinger-bridge/bridge_paths.gif +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/schrodinger-bridge/bridge_paths.png +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/schrodinger-bridge/sb_flow.gif +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/schrodinger-bridge/sb_snapshots.png +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/schrodinger-bridge/sb_trajectories.png +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/concepts/architecture.md +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/contributing.md +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/getting-started/installation.md +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/getting-started/quickstart.md +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/guides/delta-alignment.md +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/guides/flow-matching.md +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/guides/inverse-problems.md +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/guides/training.md +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/index.md +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/examples/00-foundations/01-linear-interpolant/main.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/examples/10-sampling/01-euler-flow/main.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/examples/10-sampling/02-equilibrium-matching/main.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/examples/20-training/01-flow-matching/main.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/examples/20-training/02-delta-alignment/main.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/examples/30-inverse/01-posterior/main.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/examples/90-showcase/01-landmark-viz/main.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/examples/90-showcase/02-sampling-flow-viz/main.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/examples/90-showcase/03-minibatch-ot-viz/main.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/examples/90-showcase/04-inverse-posterior-viz/main.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/examples/90-showcase/05-schrodinger-bridge-viz/main.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/examples/90-showcase/06-algorithm-comparison/main.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/examples/90-showcase/07-landmark-detection/main.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/examples/90-showcase/08-guidance-alignment-pretraining/main.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/examples/90-showcase/09-equilibrium-matching-viz/main.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/examples/90-showcase/README.md +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/mkdocs.yml +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/overrides/main.html +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/requirements-docs.txt +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/requirements.txt +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/setup.cfg +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/tests/__init__.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/tests/conftest.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/tests/test_coupling.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/tests/test_equilibrium.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/tests/test_interpolants.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/tests/test_interpolants_new.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/tests/test_inverse.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/tests/test_losses.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/tests/test_models.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/tests/test_samplers.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/tests/test_solvers.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/tests/test_trainer.py +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/torchdeltaflow.egg-info/dependency_links.txt +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/torchdeltaflow.egg-info/requires.txt +0 -0
- {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/torchdeltaflow.egg-info/top_level.txt +0 -0
|
@@ -0,0 +1,175 @@
|
|
|
1
|
+
name: publish
|
|
2
|
+
|
|
3
|
+
# Automated tagging and release for torchdeltaflow.
|
|
4
|
+
#
|
|
5
|
+
# On every merge to main this pipeline cuts a new release:
|
|
6
|
+
#
|
|
7
|
+
# 1. update-tag derives the next semantic version from the latest tag and the
|
|
8
|
+
# merge-commit message (include "#major" or "#minor" to bump those
|
|
9
|
+
# components, otherwise the patch component bumps), then creates and pushes
|
|
10
|
+
# the "vX.Y.Z" tag. The tag is the single source of truth for the version,
|
|
11
|
+
# setuptools_scm derives the package version from it (distance 0 -> clean
|
|
12
|
+
# PEP 440 release, e.g. v0.2.3 -> "0.2.3").
|
|
13
|
+
#
|
|
14
|
+
# 2. release checks out that tag, builds and validates the sdist + wheel,
|
|
15
|
+
# runs the test suite, verifies the built version matches the tag, publishes
|
|
16
|
+
# to PyPI, and opens a GitHub Release whose notes come from the matching
|
|
17
|
+
# CHANGELOG.md section (falling back to GitHub's generated notes).
|
|
18
|
+
#
|
|
19
|
+
# Because the tag is pushed with the workflow's GITHUB_TOKEN, it does not
|
|
20
|
+
# recursively trigger this workflow, so there is no separate tag-push trigger.
|
|
21
|
+
# Use "Run workflow" (workflow_dispatch) to cut a release manually.
|
|
22
|
+
#
|
|
23
|
+
# Authentication uses a PyPI API token stored as the repository secret
|
|
24
|
+
# PYPI_API_TOKEN. Create the token at https://pypi.org/manage/account/token/
|
|
25
|
+
# (scope it to the torchdeltaflow project), then add it under
|
|
26
|
+
# Settings > Secrets and variables > Actions as PYPI_API_TOKEN.
|
|
27
|
+
|
|
28
|
+
on:
|
|
29
|
+
push:
|
|
30
|
+
branches: [main]
|
|
31
|
+
workflow_dispatch:
|
|
32
|
+
|
|
33
|
+
permissions:
|
|
34
|
+
contents: write
|
|
35
|
+
|
|
36
|
+
jobs:
|
|
37
|
+
update-tag:
|
|
38
|
+
runs-on: ubuntu-latest
|
|
39
|
+
outputs:
|
|
40
|
+
new_tag: ${{ steps.create_tag.outputs.new_tag }}
|
|
41
|
+
steps:
|
|
42
|
+
- uses: actions/checkout@v4
|
|
43
|
+
with:
|
|
44
|
+
# Need the full history and all tags to find the latest version.
|
|
45
|
+
fetch-depth: 0
|
|
46
|
+
fetch-tags: true
|
|
47
|
+
|
|
48
|
+
- name: Determine version bump from commit message
|
|
49
|
+
id: bump
|
|
50
|
+
run: |
|
|
51
|
+
message=$(git log -1 --pretty=%B)
|
|
52
|
+
if echo "$message" | grep -qiE '#major'; then
|
|
53
|
+
echo "type=major" >> "$GITHUB_OUTPUT"
|
|
54
|
+
elif echo "$message" | grep -qiE '#minor'; then
|
|
55
|
+
echo "type=minor" >> "$GITHUB_OUTPUT"
|
|
56
|
+
else
|
|
57
|
+
echo "type=patch" >> "$GITHUB_OUTPUT"
|
|
58
|
+
fi
|
|
59
|
+
|
|
60
|
+
- name: Compute and push the next tag
|
|
61
|
+
id: create_tag
|
|
62
|
+
run: |
|
|
63
|
+
latest_tag=$(git tag --list 'v*' --sort=-v:refname | head -n 1)
|
|
64
|
+
bump=${{ steps.bump.outputs.type }}
|
|
65
|
+
|
|
66
|
+
if [ -z "$latest_tag" ]; then
|
|
67
|
+
new_tag="v0.1.0"
|
|
68
|
+
else
|
|
69
|
+
version=${latest_tag#v}
|
|
70
|
+
IFS='.' read -r major minor patch <<< "$version"
|
|
71
|
+
case "$bump" in
|
|
72
|
+
major) major=$((major + 1)); minor=0; patch=0 ;;
|
|
73
|
+
minor) minor=$((minor + 1)); patch=0 ;;
|
|
74
|
+
patch) patch=$((patch + 1)) ;;
|
|
75
|
+
esac
|
|
76
|
+
new_tag="v${major}.${minor}.${patch}"
|
|
77
|
+
fi
|
|
78
|
+
|
|
79
|
+
if git rev-parse "$new_tag" >/dev/null 2>&1; then
|
|
80
|
+
echo "Tag $new_tag already exists, skipping tag creation."
|
|
81
|
+
else
|
|
82
|
+
git config user.name "github-actions[bot]"
|
|
83
|
+
git config user.email "github-actions[bot]@users.noreply.github.com"
|
|
84
|
+
git tag -a "$new_tag" -m "Release $new_tag"
|
|
85
|
+
git push origin "$new_tag"
|
|
86
|
+
fi
|
|
87
|
+
|
|
88
|
+
echo "new_tag=${new_tag}" >> "$GITHUB_OUTPUT"
|
|
89
|
+
echo "Release tag: ${new_tag}"
|
|
90
|
+
|
|
91
|
+
release:
|
|
92
|
+
needs: update-tag
|
|
93
|
+
runs-on: ubuntu-latest
|
|
94
|
+
environment:
|
|
95
|
+
name: pypi
|
|
96
|
+
url: https://pypi.org/project/torchdeltaflow/
|
|
97
|
+
steps:
|
|
98
|
+
- uses: actions/checkout@v4
|
|
99
|
+
with:
|
|
100
|
+
# Check out the freshly created tag so setuptools_scm resolves the
|
|
101
|
+
# version to the tag exactly (distance 0, no ".devN").
|
|
102
|
+
ref: ${{ needs.update-tag.outputs.new_tag }}
|
|
103
|
+
fetch-depth: 0
|
|
104
|
+
fetch-tags: true
|
|
105
|
+
|
|
106
|
+
- uses: actions/setup-python@v5
|
|
107
|
+
with:
|
|
108
|
+
python-version: "3.12"
|
|
109
|
+
|
|
110
|
+
- name: Build sdist and wheel
|
|
111
|
+
run: |
|
|
112
|
+
python -m pip install --upgrade pip build
|
|
113
|
+
python -m build
|
|
114
|
+
|
|
115
|
+
- name: Check distributions
|
|
116
|
+
run: |
|
|
117
|
+
python -m pip install --upgrade twine
|
|
118
|
+
python -m twine check dist/*
|
|
119
|
+
|
|
120
|
+
- name: Verify built version matches the tag
|
|
121
|
+
run: |
|
|
122
|
+
tag="${{ needs.update-tag.outputs.new_tag }}"
|
|
123
|
+
expected="${tag#v}"
|
|
124
|
+
built=$(ls dist/*.whl | sed -E 's#.*/torchdeltaflow-([^-]+)-.*#\1#')
|
|
125
|
+
echo "Tag version: ${expected}"
|
|
126
|
+
echo "Built version: ${built}"
|
|
127
|
+
if [ "$built" != "$expected" ]; then
|
|
128
|
+
echo "::error::Built version ${built} does not match tag ${expected}."
|
|
129
|
+
exit 1
|
|
130
|
+
fi
|
|
131
|
+
|
|
132
|
+
- name: Run tests
|
|
133
|
+
run: |
|
|
134
|
+
pip install -e ".[dev]"
|
|
135
|
+
pytest --cov=deltaflow
|
|
136
|
+
|
|
137
|
+
- name: Publish to PyPI
|
|
138
|
+
uses: pypa/gh-action-pypi-publish@release/v1
|
|
139
|
+
with:
|
|
140
|
+
password: ${{ secrets.PYPI_API_TOKEN }}
|
|
141
|
+
# A manual re-run against an already-published version should be a
|
|
142
|
+
# no-op rather than a hard failure.
|
|
143
|
+
skip-existing: true
|
|
144
|
+
|
|
145
|
+
- name: Extract release notes from CHANGELOG.md
|
|
146
|
+
env:
|
|
147
|
+
RELEASE_TAG: ${{ needs.update-tag.outputs.new_tag }}
|
|
148
|
+
run: |
|
|
149
|
+
awk -v ver="${RELEASE_TAG#v}" '
|
|
150
|
+
$0 ~ ("^## \\[" ver "\\]") {found=1; next}
|
|
151
|
+
/^## \[/ {found=0}
|
|
152
|
+
found
|
|
153
|
+
' CHANGELOG.md > release_body.md
|
|
154
|
+
echo "Release notes for ${RELEASE_TAG}:"
|
|
155
|
+
cat release_body.md
|
|
156
|
+
|
|
157
|
+
- name: Create GitHub Release
|
|
158
|
+
uses: actions/github-script@v7
|
|
159
|
+
env:
|
|
160
|
+
RELEASE_TAG: ${{ needs.update-tag.outputs.new_tag }}
|
|
161
|
+
with:
|
|
162
|
+
script: |
|
|
163
|
+
const fs = require('fs');
|
|
164
|
+
let body = '';
|
|
165
|
+
try { body = fs.readFileSync('release_body.md', 'utf8').trim(); } catch {}
|
|
166
|
+
await github.rest.repos.createRelease({
|
|
167
|
+
owner: context.repo.owner,
|
|
168
|
+
repo: context.repo.repo,
|
|
169
|
+
tag_name: process.env.RELEASE_TAG,
|
|
170
|
+
name: process.env.RELEASE_TAG,
|
|
171
|
+
body: body,
|
|
172
|
+
draft: false,
|
|
173
|
+
prerelease: false,
|
|
174
|
+
generate_release_notes: true,
|
|
175
|
+
});
|
|
@@ -5,6 +5,16 @@ The format follows [Keep a Changelog](https://keepachangelog.com/en/1.1.0/).
|
|
|
5
5
|
|
|
6
6
|
## [Unreleased]
|
|
7
7
|
|
|
8
|
+
### Added
|
|
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
|
+
|
|
8
18
|
### Changed
|
|
9
19
|
|
|
10
20
|
- `scipy` is now a core dependency (previously gated behind the `ot`
|
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: torchdeltaflow
|
|
3
|
-
Version: 0.2.
|
|
3
|
+
Version: 0.2.4
|
|
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>
|
|
@@ -93,7 +93,7 @@ Dynamic: license-file
|
|
|
93
93
|
<img alt="GitHub Stars" src="https://img.shields.io/github/stars/phrugsa-limbunlom/deltaflow?style=social">
|
|
94
94
|
</a>
|
|
95
95
|
<a href="https://deepwiki.com/phrugsa-limbunlom/deltaflow" target="_blank" title="Ask DeepWiki">
|
|
96
|
-
<img alt="Ask DeepWiki" src="https://deepwiki.com/badge.svg">
|
|
96
|
+
<img alt="Ask DeepWiki" src="https://deepwiki.com/badge.svg?v=1">
|
|
97
97
|
</a>
|
|
98
98
|
<a href="https://github.com/phrugsa-limbunlom/deltaflow/actions/workflows/ci.yml" target="_blank" title="Build Status">
|
|
99
99
|
<img alt="Build Status" src="https://img.shields.io/github/actions/workflow/status/phrugsa-limbunlom/deltaflow/ci.yml?branch=main&style=flat-square&label=build&color=3f9e73">
|
|
@@ -16,7 +16,7 @@
|
|
|
16
16
|
<img alt="GitHub Stars" src="https://img.shields.io/github/stars/phrugsa-limbunlom/deltaflow?style=social">
|
|
17
17
|
</a>
|
|
18
18
|
<a href="https://deepwiki.com/phrugsa-limbunlom/deltaflow" target="_blank" title="Ask DeepWiki">
|
|
19
|
-
<img alt="Ask DeepWiki" src="https://deepwiki.com/badge.svg">
|
|
19
|
+
<img alt="Ask DeepWiki" src="https://deepwiki.com/badge.svg?v=1">
|
|
20
20
|
</a>
|
|
21
21
|
<a href="https://github.com/phrugsa-limbunlom/deltaflow/actions/workflows/ci.yml" target="_blank" title="Build Status">
|
|
22
22
|
<img alt="Build Status" src="https://img.shields.io/github/actions/workflow/status/phrugsa-limbunlom/deltaflow/ci.yml?branch=main&style=flat-square&label=build&color=3f9e73">
|
|
@@ -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
|
]
|
|
@@ -0,0 +1,438 @@
|
|
|
1
|
+
"""Diffusion-Transformer (DiT) velocity field with adaLN-Zero conditioning.
|
|
2
|
+
|
|
3
|
+
This module provides a native transformer backbone for DeltaFlow, conditioned
|
|
4
|
+
through **Adaptive Layer Normalization with Zero initialization (adaLN-Zero)**.
|
|
5
|
+
The displacement/transport framing still holds: the network regresses the
|
|
6
|
+
flow-matching target velocity ``Delta = x1 - x0`` (see
|
|
7
|
+
`deltaflow.losses.ConditionalFlowMatchingLoss`), only now the field is a
|
|
8
|
+
sequence model over image patches rather than a convolutional UNet.
|
|
9
|
+
|
|
10
|
+
The conditioning path follows Peebles & Xie (2023). A shared conditioning
|
|
11
|
+
embedding ``c`` (time embedding plus, optionally, a class embedding) drives a
|
|
12
|
+
small MLP per block that emits per-channel ``(shift, scale, gate)`` modulation
|
|
13
|
+
parameters. The gates, and the final output projection, are zero-initialized, so
|
|
14
|
+
every residual branch starts as the identity and the whole field starts at
|
|
15
|
+
``v = 0``. Training then departs smoothly from that stable fixed point, which is
|
|
16
|
+
the stability trick that makes deep DiTs trainable without warmup tricks.
|
|
17
|
+
|
|
18
|
+
Classifier-free guidance (CFG) is supported through a learned *null* class token
|
|
19
|
+
in `LabelEmbedding`: at train time a fraction of labels are dropped to the null
|
|
20
|
+
token, and at sample time `DiT.forward_with_cfg` extrapolates between the
|
|
21
|
+
conditional and unconditional velocity.
|
|
22
|
+
|
|
23
|
+
References:
|
|
24
|
+
Peebles & Xie, "Scalable Diffusion Models with Transformers" (2023),
|
|
25
|
+
https://arxiv.org/abs/2212.09748.
|
|
26
|
+
Ho & Salimans, "Classifier-Free Diffusion Guidance" (2022),
|
|
27
|
+
https://arxiv.org/abs/2207.12598.
|
|
28
|
+
"""
|
|
29
|
+
|
|
30
|
+
import math
|
|
31
|
+
from typing import Optional
|
|
32
|
+
|
|
33
|
+
import torch
|
|
34
|
+
import torch.nn as nn
|
|
35
|
+
import torch.nn.functional as F
|
|
36
|
+
|
|
37
|
+
from ..core.base_velocity_field import BaseVelocityField
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
def modulate(x: torch.Tensor, shift: torch.Tensor, scale: torch.Tensor) -> torch.Tensor:
|
|
41
|
+
"""Apply adaLN modulation ``x * (1 + scale) + shift`` with broadcasting.
|
|
42
|
+
|
|
43
|
+
``x`` has shape ``(B, N, D)`` (batch, tokens, channels), while ``shift`` and
|
|
44
|
+
``scale`` have shape ``(B, D)``. The ``1 +`` keeps the transform centred on
|
|
45
|
+
the identity, so a zero-initialized modulation MLP leaves ``x`` unchanged.
|
|
46
|
+
"""
|
|
47
|
+
return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1)
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
class TimestepEmbedding(nn.Module):
|
|
51
|
+
"""Reusable sinusoidal timestep embedding followed by a small MLP.
|
|
52
|
+
|
|
53
|
+
The continuous flow-matching time ``t in [0, 1]`` (DeltaFlow convention:
|
|
54
|
+
``t=0`` is noise, ``t=1`` is data) is first multiplied by ``time_scale`` to
|
|
55
|
+
spread it across the sinusoidal frequency band, then embedded with the usual
|
|
56
|
+
transformer sinusoids and projected by a two-layer MLP.
|
|
57
|
+
|
|
58
|
+
Args:
|
|
59
|
+
hidden_size: output (and MLP) width.
|
|
60
|
+
frequency_dim: width of the raw sinusoidal features before the MLP.
|
|
61
|
+
time_scale: multiplier applied to ``t`` before the sinusoids. The
|
|
62
|
+
default of ``1000`` mirrors the diffusion-timestep range DiT was
|
|
63
|
+
tuned on and gives continuous ``t in [0, 1]`` enough resolution.
|
|
64
|
+
max_period: controls the lowest sinusoidal frequency.
|
|
65
|
+
"""
|
|
66
|
+
|
|
67
|
+
def __init__(
|
|
68
|
+
self,
|
|
69
|
+
hidden_size: int,
|
|
70
|
+
frequency_dim: int = 256,
|
|
71
|
+
time_scale: float = 1000.0,
|
|
72
|
+
max_period: int = 10000,
|
|
73
|
+
):
|
|
74
|
+
super().__init__()
|
|
75
|
+
self.frequency_dim = frequency_dim
|
|
76
|
+
self.time_scale = time_scale
|
|
77
|
+
self.max_period = max_period
|
|
78
|
+
self.mlp = nn.Sequential(
|
|
79
|
+
nn.Linear(frequency_dim, hidden_size),
|
|
80
|
+
nn.SiLU(),
|
|
81
|
+
nn.Linear(hidden_size, hidden_size),
|
|
82
|
+
)
|
|
83
|
+
|
|
84
|
+
def _sinusoidal(self, t: torch.Tensor) -> torch.Tensor:
|
|
85
|
+
half = self.frequency_dim // 2
|
|
86
|
+
freqs = torch.exp(
|
|
87
|
+
-math.log(self.max_period)
|
|
88
|
+
* torch.arange(half, dtype=torch.float32, device=t.device)
|
|
89
|
+
/ half
|
|
90
|
+
)
|
|
91
|
+
args = t[:, None].float() * freqs[None]
|
|
92
|
+
emb = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
|
|
93
|
+
if self.frequency_dim % 2:
|
|
94
|
+
emb = torch.cat([emb, torch.zeros_like(emb[:, :1])], dim=-1)
|
|
95
|
+
return emb
|
|
96
|
+
|
|
97
|
+
def forward(self, t: torch.Tensor) -> torch.Tensor:
|
|
98
|
+
if t.dim() == 0:
|
|
99
|
+
t = t.expand(1)
|
|
100
|
+
emb = self._sinusoidal(t * self.time_scale)
|
|
101
|
+
return self.mlp(emb.to(self.mlp[0].weight.dtype))
|
|
102
|
+
|
|
103
|
+
|
|
104
|
+
class LabelEmbedding(nn.Module):
|
|
105
|
+
"""Class-label embedding with a learned null token for guidance.
|
|
106
|
+
|
|
107
|
+
The embedding table has ``num_classes + 1`` rows; the extra row is the
|
|
108
|
+
*null* (unconditional) token used by classifier-free guidance. During
|
|
109
|
+
training a fraction ``dropout_prob`` of labels are replaced by the null
|
|
110
|
+
token, teaching the field both the conditional and unconditional velocity
|
|
111
|
+
with one set of weights.
|
|
112
|
+
|
|
113
|
+
Args:
|
|
114
|
+
num_classes: number of real classes.
|
|
115
|
+
hidden_size: embedding width.
|
|
116
|
+
dropout_prob: probability of dropping a label to the null token at
|
|
117
|
+
train time (set ``0`` to disable CFG training).
|
|
118
|
+
"""
|
|
119
|
+
|
|
120
|
+
def __init__(self, num_classes: int, hidden_size: int, dropout_prob: float = 0.1):
|
|
121
|
+
super().__init__()
|
|
122
|
+
self.num_classes = num_classes
|
|
123
|
+
self.dropout_prob = dropout_prob
|
|
124
|
+
self.null_index = num_classes
|
|
125
|
+
self.embedding_table = nn.Embedding(num_classes + 1, hidden_size)
|
|
126
|
+
|
|
127
|
+
def token_drop(
|
|
128
|
+
self, labels: torch.Tensor, force_drop_ids: Optional[torch.Tensor] = None
|
|
129
|
+
) -> torch.Tensor:
|
|
130
|
+
"""Replace a random subset of ``labels`` with the null token.
|
|
131
|
+
|
|
132
|
+
If ``force_drop_ids`` is given (a boolean/0-1 mask), those positions are
|
|
133
|
+
dropped deterministically instead, which is how the unconditional branch
|
|
134
|
+
is requested at sampling time.
|
|
135
|
+
"""
|
|
136
|
+
if force_drop_ids is None:
|
|
137
|
+
drop = torch.rand(labels.shape[0], device=labels.device) < self.dropout_prob
|
|
138
|
+
else:
|
|
139
|
+
drop = force_drop_ids.to(torch.bool)
|
|
140
|
+
return torch.where(drop, torch.full_like(labels, self.null_index), labels)
|
|
141
|
+
|
|
142
|
+
def forward(
|
|
143
|
+
self,
|
|
144
|
+
labels: torch.Tensor,
|
|
145
|
+
train: Optional[bool] = None,
|
|
146
|
+
force_drop_ids: Optional[torch.Tensor] = None,
|
|
147
|
+
) -> torch.Tensor:
|
|
148
|
+
train = self.training if train is None else train
|
|
149
|
+
if (train and self.dropout_prob > 0) or force_drop_ids is not None:
|
|
150
|
+
labels = self.token_drop(labels, force_drop_ids)
|
|
151
|
+
return self.embedding_table(labels)
|
|
152
|
+
|
|
153
|
+
|
|
154
|
+
class _Attention(nn.Module):
|
|
155
|
+
"""Minimal multi-head self-attention over token sequences ``(B, N, D)``."""
|
|
156
|
+
|
|
157
|
+
def __init__(self, hidden_size: int, num_heads: int):
|
|
158
|
+
super().__init__()
|
|
159
|
+
if hidden_size % num_heads != 0:
|
|
160
|
+
raise ValueError(
|
|
161
|
+
f"hidden_size ({hidden_size}) must be divisible by num_heads ({num_heads})"
|
|
162
|
+
)
|
|
163
|
+
self.num_heads = num_heads
|
|
164
|
+
self.head_dim = hidden_size // num_heads
|
|
165
|
+
self.qkv = nn.Linear(hidden_size, hidden_size * 3)
|
|
166
|
+
self.proj = nn.Linear(hidden_size, hidden_size)
|
|
167
|
+
|
|
168
|
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
|
169
|
+
B, N, D = x.shape
|
|
170
|
+
qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, self.head_dim)
|
|
171
|
+
qkv = qkv.permute(2, 0, 3, 1, 4) # (3, B, heads, N, head_dim)
|
|
172
|
+
q, k, v = qkv[0], qkv[1], qkv[2]
|
|
173
|
+
out = F.scaled_dot_product_attention(q, k, v)
|
|
174
|
+
out = out.transpose(1, 2).reshape(B, N, D)
|
|
175
|
+
return self.proj(out)
|
|
176
|
+
|
|
177
|
+
|
|
178
|
+
class DiTBlock(nn.Module):
|
|
179
|
+
"""A transformer block modulated by adaLN-Zero conditioning.
|
|
180
|
+
|
|
181
|
+
Both sub-layers (self-attention and MLP) are wrapped as
|
|
182
|
+
|
|
183
|
+
``x = x + gate * sublayer(modulate(norm(x), shift, scale))``,
|
|
184
|
+
|
|
185
|
+
where ``(shift, scale, gate)`` are produced per sub-layer from the shared
|
|
186
|
+
conditioning embedding ``c``. The six modulation vectors come from a single
|
|
187
|
+
``SiLU -> Linear`` head whose weights are zero-initialized (see
|
|
188
|
+
`DiT.initialize_weights`), so at initialization every gate is zero and the
|
|
189
|
+
block is the identity map.
|
|
190
|
+
|
|
191
|
+
Args:
|
|
192
|
+
hidden_size: token channel width.
|
|
193
|
+
num_heads: attention heads.
|
|
194
|
+
mlp_ratio: hidden expansion of the feed-forward MLP.
|
|
195
|
+
"""
|
|
196
|
+
|
|
197
|
+
def __init__(self, hidden_size: int, num_heads: int, mlp_ratio: float = 4.0):
|
|
198
|
+
super().__init__()
|
|
199
|
+
self.norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
|
200
|
+
self.attn = _Attention(hidden_size, num_heads)
|
|
201
|
+
self.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
|
202
|
+
mlp_hidden = int(hidden_size * mlp_ratio)
|
|
203
|
+
self.mlp = nn.Sequential(
|
|
204
|
+
nn.Linear(hidden_size, mlp_hidden),
|
|
205
|
+
nn.GELU(approximate="tanh"),
|
|
206
|
+
nn.Linear(mlp_hidden, hidden_size),
|
|
207
|
+
)
|
|
208
|
+
self.adaLN_modulation = nn.Sequential(
|
|
209
|
+
nn.SiLU(),
|
|
210
|
+
nn.Linear(hidden_size, 6 * hidden_size),
|
|
211
|
+
)
|
|
212
|
+
|
|
213
|
+
def forward(self, x: torch.Tensor, c: torch.Tensor) -> torch.Tensor:
|
|
214
|
+
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.adaLN_modulation(
|
|
215
|
+
c
|
|
216
|
+
).chunk(6, dim=-1)
|
|
217
|
+
x = x + gate_msa.unsqueeze(1) * self.attn(modulate(self.norm1(x), shift_msa, scale_msa))
|
|
218
|
+
x = x + gate_mlp.unsqueeze(1) * self.mlp(modulate(self.norm2(x), shift_mlp, scale_mlp))
|
|
219
|
+
return x
|
|
220
|
+
|
|
221
|
+
|
|
222
|
+
class FinalLayer(nn.Module):
|
|
223
|
+
"""adaLN-Zero output head that maps tokens back to patch pixels.
|
|
224
|
+
|
|
225
|
+
The final normalization is modulated by ``(shift, scale)`` from the
|
|
226
|
+
conditioning embedding, then a linear layer projects each token to
|
|
227
|
+
``patch_size**2 * out_channels`` values. Both the modulation head and the
|
|
228
|
+
output projection are zero-initialized so the field predicts ``v = 0`` at the
|
|
229
|
+
start of training.
|
|
230
|
+
"""
|
|
231
|
+
|
|
232
|
+
def __init__(self, hidden_size: int, patch_size: int, out_channels: int):
|
|
233
|
+
super().__init__()
|
|
234
|
+
self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
|
235
|
+
self.linear = nn.Linear(hidden_size, patch_size * patch_size * out_channels)
|
|
236
|
+
self.adaLN_modulation = nn.Sequential(
|
|
237
|
+
nn.SiLU(),
|
|
238
|
+
nn.Linear(hidden_size, 2 * hidden_size),
|
|
239
|
+
)
|
|
240
|
+
|
|
241
|
+
def forward(self, x: torch.Tensor, c: torch.Tensor) -> torch.Tensor:
|
|
242
|
+
shift, scale = self.adaLN_modulation(c).chunk(2, dim=-1)
|
|
243
|
+
x = modulate(self.norm_final(x), shift, scale)
|
|
244
|
+
return self.linear(x)
|
|
245
|
+
|
|
246
|
+
|
|
247
|
+
class DiT(BaseVelocityField):
|
|
248
|
+
r"""A class-conditional Diffusion Transformer velocity field.
|
|
249
|
+
|
|
250
|
+
The field patchifies an image ``x`` of shape ``(B, C, H, W)`` into a token
|
|
251
|
+
sequence, adds a learned position embedding, runs it through ``depth``
|
|
252
|
+
`DiTBlock` layers conditioned on ``time + class`` embeddings, and unpatchifies
|
|
253
|
+
the output back to a velocity of the same shape as ``x``. It satisfies the
|
|
254
|
+
`BaseVelocityField` contract, so it drops straight into DeltaFlow's losses
|
|
255
|
+
and solvers.
|
|
256
|
+
|
|
257
|
+
Conditioning is passed through the opaque ``**cond`` channel as ``cond["y"]``
|
|
258
|
+
(integer class labels of shape ``(B,)``). When ``y`` is omitted the field
|
|
259
|
+
runs unconditionally via the learned null token.
|
|
260
|
+
|
|
261
|
+
Args:
|
|
262
|
+
input_size: spatial size of the (square) input image.
|
|
263
|
+
patch_size: side length of each square patch; must divide ``input_size``.
|
|
264
|
+
in_channels: number of image channels.
|
|
265
|
+
hidden_size: transformer token width.
|
|
266
|
+
depth: number of `DiTBlock` layers.
|
|
267
|
+
num_heads: attention heads per block.
|
|
268
|
+
mlp_ratio: feed-forward expansion ratio.
|
|
269
|
+
num_classes: number of conditioning classes (a null token is added on
|
|
270
|
+
top for classifier-free guidance).
|
|
271
|
+
class_dropout_prob: train-time label-dropout probability for CFG.
|
|
272
|
+
time_scale: multiplier applied to ``t`` inside `TimestepEmbedding`.
|
|
273
|
+
"""
|
|
274
|
+
|
|
275
|
+
def __init__(
|
|
276
|
+
self,
|
|
277
|
+
input_size: int = 32,
|
|
278
|
+
patch_size: int = 4,
|
|
279
|
+
in_channels: int = 1,
|
|
280
|
+
hidden_size: int = 256,
|
|
281
|
+
depth: int = 4,
|
|
282
|
+
num_heads: int = 4,
|
|
283
|
+
mlp_ratio: float = 4.0,
|
|
284
|
+
num_classes: int = 10,
|
|
285
|
+
class_dropout_prob: float = 0.1,
|
|
286
|
+
time_scale: float = 1000.0,
|
|
287
|
+
):
|
|
288
|
+
super().__init__()
|
|
289
|
+
if input_size % patch_size != 0:
|
|
290
|
+
raise ValueError(
|
|
291
|
+
f"input_size ({input_size}) must be divisible by patch_size ({patch_size})"
|
|
292
|
+
)
|
|
293
|
+
self.in_channels = in_channels
|
|
294
|
+
self.out_channels = in_channels
|
|
295
|
+
self.patch_size = patch_size
|
|
296
|
+
self.input_size = input_size
|
|
297
|
+
self.num_classes = num_classes
|
|
298
|
+
self.num_patches_side = input_size // patch_size
|
|
299
|
+
self.num_patches = self.num_patches_side**2
|
|
300
|
+
|
|
301
|
+
self.patch_embed = nn.Conv2d(
|
|
302
|
+
in_channels, hidden_size, kernel_size=patch_size, stride=patch_size
|
|
303
|
+
)
|
|
304
|
+
self.pos_embed = nn.Parameter(torch.zeros(1, self.num_patches, hidden_size))
|
|
305
|
+
self.t_embedder = TimestepEmbedding(hidden_size, time_scale=time_scale)
|
|
306
|
+
self.y_embedder = LabelEmbedding(num_classes, hidden_size, class_dropout_prob)
|
|
307
|
+
|
|
308
|
+
self.blocks = nn.ModuleList(
|
|
309
|
+
[DiTBlock(hidden_size, num_heads, mlp_ratio) for _ in range(depth)]
|
|
310
|
+
)
|
|
311
|
+
self.final_layer = FinalLayer(hidden_size, patch_size, self.out_channels)
|
|
312
|
+
|
|
313
|
+
self.initialize_weights()
|
|
314
|
+
|
|
315
|
+
def initialize_weights(self) -> None:
|
|
316
|
+
"""Xavier-init linear layers, then zero-init every adaLN-Zero gate."""
|
|
317
|
+
|
|
318
|
+
def _basic_init(module: nn.Module) -> None:
|
|
319
|
+
if isinstance(module, nn.Linear):
|
|
320
|
+
nn.init.xavier_uniform_(module.weight)
|
|
321
|
+
if module.bias is not None:
|
|
322
|
+
nn.init.zeros_(module.bias)
|
|
323
|
+
|
|
324
|
+
self.apply(_basic_init)
|
|
325
|
+
|
|
326
|
+
nn.init.normal_(self.pos_embed, std=0.02)
|
|
327
|
+
|
|
328
|
+
w = self.patch_embed.weight.data
|
|
329
|
+
nn.init.xavier_uniform_(w.view(w.shape[0], -1))
|
|
330
|
+
nn.init.zeros_(self.patch_embed.bias)
|
|
331
|
+
|
|
332
|
+
nn.init.normal_(self.y_embedder.embedding_table.weight, std=0.02)
|
|
333
|
+
nn.init.normal_(self.t_embedder.mlp[0].weight, std=0.02)
|
|
334
|
+
nn.init.normal_(self.t_embedder.mlp[2].weight, std=0.02)
|
|
335
|
+
|
|
336
|
+
# adaLN-Zero: zero the modulation heads so blocks start as the identity.
|
|
337
|
+
for block in self.blocks:
|
|
338
|
+
nn.init.zeros_(block.adaLN_modulation[-1].weight)
|
|
339
|
+
nn.init.zeros_(block.adaLN_modulation[-1].bias)
|
|
340
|
+
|
|
341
|
+
# Zero the final layer so the initial predicted velocity is exactly 0.
|
|
342
|
+
nn.init.zeros_(self.final_layer.adaLN_modulation[-1].weight)
|
|
343
|
+
nn.init.zeros_(self.final_layer.adaLN_modulation[-1].bias)
|
|
344
|
+
nn.init.zeros_(self.final_layer.linear.weight)
|
|
345
|
+
nn.init.zeros_(self.final_layer.linear.bias)
|
|
346
|
+
|
|
347
|
+
def _unpatchify(self, x: torch.Tensor) -> torch.Tensor:
|
|
348
|
+
"""``(B, num_patches, p*p*C) -> (B, C, H, W)``."""
|
|
349
|
+
c = self.out_channels
|
|
350
|
+
p = self.patch_size
|
|
351
|
+
n = self.num_patches_side
|
|
352
|
+
x = x.reshape(x.shape[0], n, n, p, p, c)
|
|
353
|
+
x = torch.einsum("bhwpqc->bchpwq", x)
|
|
354
|
+
return x.reshape(x.shape[0], c, n * p, n * p)
|
|
355
|
+
|
|
356
|
+
def _conditioning(
|
|
357
|
+
self,
|
|
358
|
+
x: torch.Tensor,
|
|
359
|
+
t: torch.Tensor,
|
|
360
|
+
y: Optional[torch.Tensor],
|
|
361
|
+
force_drop_ids: Optional[torch.Tensor],
|
|
362
|
+
) -> torch.Tensor:
|
|
363
|
+
c = self.t_embedder(t)
|
|
364
|
+
if y is None:
|
|
365
|
+
y = torch.full((x.shape[0],), self.y_embedder.null_index, device=x.device)
|
|
366
|
+
c = c + self.y_embedder(y, train=False)
|
|
367
|
+
else:
|
|
368
|
+
c = c + self.y_embedder(y.to(x.device), force_drop_ids=force_drop_ids)
|
|
369
|
+
return c
|
|
370
|
+
|
|
371
|
+
def forward(
|
|
372
|
+
self,
|
|
373
|
+
x: torch.Tensor,
|
|
374
|
+
t: torch.Tensor,
|
|
375
|
+
y: Optional[torch.Tensor] = None,
|
|
376
|
+
force_drop_ids: Optional[torch.Tensor] = None,
|
|
377
|
+
**cond,
|
|
378
|
+
) -> torch.Tensor:
|
|
379
|
+
"""Predict the velocity ``v_theta(x, t, y)``.
|
|
380
|
+
|
|
381
|
+
Args:
|
|
382
|
+
x: input image batch ``(B, C, H, W)``.
|
|
383
|
+
t: time of shape ``(B,)`` or a scalar (DeltaFlow convention
|
|
384
|
+
``t=0`` noise, ``t=1`` data).
|
|
385
|
+
y: optional integer class labels ``(B,)`` forwarded as ``cond["y"]``;
|
|
386
|
+
omit for unconditional inference.
|
|
387
|
+
force_drop_ids: optional mask selecting which labels to force to the
|
|
388
|
+
null token (used by the unconditional branch of CFG).
|
|
389
|
+
"""
|
|
390
|
+
if not isinstance(t, torch.Tensor):
|
|
391
|
+
t = torch.tensor(t, device=x.device)
|
|
392
|
+
if t.dim() == 0:
|
|
393
|
+
t = t.expand(x.shape[0])
|
|
394
|
+
|
|
395
|
+
h = self.patch_embed(x).flatten(2).transpose(1, 2) # (B, N, D)
|
|
396
|
+
h = h + self.pos_embed
|
|
397
|
+
c = self._conditioning(x, t, y, force_drop_ids)
|
|
398
|
+
for block in self.blocks:
|
|
399
|
+
h = block(h, c)
|
|
400
|
+
h = self.final_layer(h, c)
|
|
401
|
+
return self._unpatchify(h)
|
|
402
|
+
|
|
403
|
+
@torch.no_grad()
|
|
404
|
+
def forward_with_cfg(
|
|
405
|
+
self, x: torch.Tensor, t: torch.Tensor, y: torch.Tensor, cfg_scale: float = 4.0
|
|
406
|
+
) -> torch.Tensor:
|
|
407
|
+
r"""Classifier-free-guided velocity for sampling.
|
|
408
|
+
|
|
409
|
+
Runs the field once with the real labels and once with the null token,
|
|
410
|
+
then extrapolates
|
|
411
|
+
|
|
412
|
+
\[
|
|
413
|
+
v_{\mathrm{cfg}} = v_\text{uncond}
|
|
414
|
+
+ s\,\bigl(v_\text{cond} - v_\text{uncond}\bigr),
|
|
415
|
+
\]
|
|
416
|
+
|
|
417
|
+
with guidance weight ``s = cfg_scale``. The two passes are batched
|
|
418
|
+
together, so this costs one forward over a doubled batch.
|
|
419
|
+
"""
|
|
420
|
+
half = x
|
|
421
|
+
combined = torch.cat([half, half], dim=0)
|
|
422
|
+
t_cat = torch.cat([t, t], dim=0) if t.dim() > 0 else t
|
|
423
|
+
y_cat = torch.cat([y, y], dim=0)
|
|
424
|
+
drop = torch.zeros(y_cat.shape[0], dtype=torch.bool, device=y.device)
|
|
425
|
+
drop[y.shape[0] :] = True
|
|
426
|
+
v = self.forward(combined, t_cat, y=y_cat, force_drop_ids=drop)
|
|
427
|
+
v_cond, v_uncond = v.chunk(2, dim=0)
|
|
428
|
+
return v_uncond + cfg_scale * (v_cond - v_uncond)
|
|
429
|
+
|
|
430
|
+
|
|
431
|
+
__all__ = [
|
|
432
|
+
"DiT",
|
|
433
|
+
"DiTBlock",
|
|
434
|
+
"FinalLayer",
|
|
435
|
+
"LabelEmbedding",
|
|
436
|
+
"TimestepEmbedding",
|
|
437
|
+
"modulate",
|
|
438
|
+
]
|