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.
Files changed (153) hide show
  1. torchdeltaflow-0.2.4/.github/workflows/publish.yml +175 -0
  2. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/CHANGELOG.md +10 -0
  3. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/PKG-INFO +2 -2
  4. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/README.md +1 -1
  5. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/models/__init__.py +9 -1
  6. torchdeltaflow-0.2.4/deltaflow/models/dit.py +438 -0
  7. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/examples.md +36 -0
  8. torchdeltaflow-0.2.4/examples/20-training/03-conditional-dit/main.py +72 -0
  9. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/pyproject.toml +5 -4
  10. torchdeltaflow-0.2.4/tests/test_dit.py +147 -0
  11. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/torchdeltaflow.egg-info/PKG-INFO +2 -2
  12. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/torchdeltaflow.egg-info/SOURCES.txt +3 -0
  13. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/torchdeltaflow.egg-info/scm_file_list.json +3 -0
  14. torchdeltaflow-0.2.4/torchdeltaflow.egg-info/scm_version.json +8 -0
  15. torchdeltaflow-0.2.3.dev21/.github/workflows/publish.yml +0 -78
  16. torchdeltaflow-0.2.3.dev21/torchdeltaflow.egg-info/scm_version.json +0 -8
  17. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/.github/workflows/ci.yml +0 -0
  18. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/.github/workflows/docs.yml +0 -0
  19. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/.gitignore +0 -0
  20. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/CITATION.cff +0 -0
  21. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/CODE_OF_CONDUCT.md +0 -0
  22. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/CONTRIBUTING.md +0 -0
  23. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/LICENSE +0 -0
  24. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/MANIFEST.in +0 -0
  25. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/benchmarks/README.md +0 -0
  26. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/__init__.py +0 -0
  27. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/core/__init__.py +0 -0
  28. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/core/base.py +0 -0
  29. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/core/base_coupling.py +0 -0
  30. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/core/base_equilibrium_field.py +0 -0
  31. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/core/base_equilibrium_interpolant.py +0 -0
  32. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/core/base_interpolant.py +0 -0
  33. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/core/base_loss.py +0 -0
  34. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/core/base_solver.py +0 -0
  35. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/core/base_velocity_field.py +0 -0
  36. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/datasets/__init__.py +0 -0
  37. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/datasets/radiograph.py +0 -0
  38. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/interpolants/__init__.py +0 -0
  39. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/interpolants/base.py +0 -0
  40. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/interpolants/equilibrium.py +0 -0
  41. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/interpolants/linear.py +0 -0
  42. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/interpolants/ot.py +0 -0
  43. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/interpolants/schrodinger_bridge.py +0 -0
  44. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/interpolants/variance_preserving.py +0 -0
  45. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/inverse/__init__.py +0 -0
  46. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/inverse/likelihood.py +0 -0
  47. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/inverse/operators.py +0 -0
  48. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/inverse/tweedie.py +0 -0
  49. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/losses/__init__.py +0 -0
  50. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/losses/conditional_flow_matching.py +0 -0
  51. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/losses/delta_alignment.py +0 -0
  52. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/losses/equilibrium_matching.py +0 -0
  53. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/losses/flow_matching.py +0 -0
  54. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/models/backbone.py +0 -0
  55. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/models/ema.py +0 -0
  56. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/models/projector.py +0 -0
  57. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/samplers/__init__.py +0 -0
  58. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/samplers/euler.py +0 -0
  59. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/solvers/__init__.py +0 -0
  60. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/solvers/euler.py +0 -0
  61. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/solvers/gradient_descent.py +0 -0
  62. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/solvers/heun.py +0 -0
  63. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/solvers/posterior_solver.py +0 -0
  64. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/trainer/__init__.py +0 -0
  65. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/trainer/coupling.py +0 -0
  66. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/trainer/data.py +0 -0
  67. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/trainer/loop.py +0 -0
  68. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/utils/__init__.py +0 -0
  69. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/utils/numerical.py +0 -0
  70. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/deltaflow/utils/ot.py +0 -0
  71. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/api/core.md +0 -0
  72. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/api/datasets.md +0 -0
  73. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/api/index.md +0 -0
  74. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/api/interpolants.md +0 -0
  75. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/api/inverse.md +0 -0
  76. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/api/losses.md +0 -0
  77. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/api/models.md +0 -0
  78. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/api/samplers.md +0 -0
  79. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/api/solvers.md +0 -0
  80. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/api/trainer.md +0 -0
  81. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/api/utils.md +0 -0
  82. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/algorithm-comparison/comparison.gif +0 -0
  83. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/algorithm-comparison/comparison.png +0 -0
  84. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/algorithm-comparison/trajectories_comparison.png +0 -0
  85. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/css/extra.css +0 -0
  86. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/equilibrium-matching/energy_landscape.png +0 -0
  87. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/equilibrium-matching/eqm_sampling.gif +0 -0
  88. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/equilibrium-matching/gd_snapshots.png +0 -0
  89. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/equilibrium-matching/gd_trajectories.png +0 -0
  90. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/favicon.ico +0 -0
  91. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/favicon.png +0 -0
  92. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/guidance-alignment/features.png +0 -0
  93. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/inverse-posterior/inverse_posterior.gif +0 -0
  94. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/inverse-posterior/inverse_posterior.png +0 -0
  95. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/js/mathjax.js +0 -0
  96. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/landmark-detection/landmark_detection.png +0 -0
  97. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/logo.svg +0 -0
  98. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/minibatch-ot/minibatch_ot.gif +0 -0
  99. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/minibatch-ot/minibatch_ot.png +0 -0
  100. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/sampling-flow/flow.gif +0 -0
  101. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/sampling-flow/snapshots.png +0 -0
  102. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/sampling-flow/trajectories.png +0 -0
  103. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/sampling-flow/velocity_field.png +0 -0
  104. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/schrodinger-bridge/bridge_paths.gif +0 -0
  105. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/schrodinger-bridge/bridge_paths.png +0 -0
  106. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/schrodinger-bridge/sb_flow.gif +0 -0
  107. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/schrodinger-bridge/sb_snapshots.png +0 -0
  108. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/assets/schrodinger-bridge/sb_trajectories.png +0 -0
  109. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/concepts/architecture.md +0 -0
  110. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/contributing.md +0 -0
  111. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/getting-started/installation.md +0 -0
  112. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/getting-started/quickstart.md +0 -0
  113. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/guides/delta-alignment.md +0 -0
  114. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/guides/flow-matching.md +0 -0
  115. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/guides/inverse-problems.md +0 -0
  116. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/guides/training.md +0 -0
  117. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/docs/index.md +0 -0
  118. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/examples/00-foundations/01-linear-interpolant/main.py +0 -0
  119. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/examples/10-sampling/01-euler-flow/main.py +0 -0
  120. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/examples/10-sampling/02-equilibrium-matching/main.py +0 -0
  121. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/examples/20-training/01-flow-matching/main.py +0 -0
  122. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/examples/20-training/02-delta-alignment/main.py +0 -0
  123. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/examples/30-inverse/01-posterior/main.py +0 -0
  124. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/examples/90-showcase/01-landmark-viz/main.py +0 -0
  125. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/examples/90-showcase/02-sampling-flow-viz/main.py +0 -0
  126. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/examples/90-showcase/03-minibatch-ot-viz/main.py +0 -0
  127. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/examples/90-showcase/04-inverse-posterior-viz/main.py +0 -0
  128. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/examples/90-showcase/05-schrodinger-bridge-viz/main.py +0 -0
  129. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/examples/90-showcase/06-algorithm-comparison/main.py +0 -0
  130. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/examples/90-showcase/07-landmark-detection/main.py +0 -0
  131. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/examples/90-showcase/08-guidance-alignment-pretraining/main.py +0 -0
  132. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/examples/90-showcase/09-equilibrium-matching-viz/main.py +0 -0
  133. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/examples/90-showcase/README.md +0 -0
  134. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/mkdocs.yml +0 -0
  135. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/overrides/main.html +0 -0
  136. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/requirements-docs.txt +0 -0
  137. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/requirements.txt +0 -0
  138. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/setup.cfg +0 -0
  139. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/tests/__init__.py +0 -0
  140. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/tests/conftest.py +0 -0
  141. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/tests/test_coupling.py +0 -0
  142. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/tests/test_equilibrium.py +0 -0
  143. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/tests/test_interpolants.py +0 -0
  144. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/tests/test_interpolants_new.py +0 -0
  145. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/tests/test_inverse.py +0 -0
  146. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/tests/test_losses.py +0 -0
  147. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/tests/test_models.py +0 -0
  148. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/tests/test_samplers.py +0 -0
  149. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/tests/test_solvers.py +0 -0
  150. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/tests/test_trainer.py +0 -0
  151. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/torchdeltaflow.egg-info/dependency_links.txt +0 -0
  152. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.4}/torchdeltaflow.egg-info/requires.txt +0 -0
  153. {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.dev21
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
+ ]