torchdeltaflow 0.2.3.dev21__tar.gz → 0.2.3.dev23__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (152) hide show
  1. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/CHANGELOG.md +10 -0
  2. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/PKG-INFO +1 -1
  3. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/models/__init__.py +9 -1
  4. torchdeltaflow-0.2.3.dev23/deltaflow/models/dit.py +438 -0
  5. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/examples.md +36 -0
  6. torchdeltaflow-0.2.3.dev23/examples/20-training/03-conditional-dit/main.py +72 -0
  7. torchdeltaflow-0.2.3.dev23/tests/test_dit.py +147 -0
  8. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/torchdeltaflow.egg-info/PKG-INFO +1 -1
  9. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/torchdeltaflow.egg-info/SOURCES.txt +3 -0
  10. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/torchdeltaflow.egg-info/scm_file_list.json +3 -0
  11. torchdeltaflow-0.2.3.dev23/torchdeltaflow.egg-info/scm_version.json +8 -0
  12. torchdeltaflow-0.2.3.dev21/torchdeltaflow.egg-info/scm_version.json +0 -8
  13. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/.github/workflows/ci.yml +0 -0
  14. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/.github/workflows/docs.yml +0 -0
  15. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/.github/workflows/publish.yml +0 -0
  16. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/.gitignore +0 -0
  17. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/CITATION.cff +0 -0
  18. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/CODE_OF_CONDUCT.md +0 -0
  19. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/CONTRIBUTING.md +0 -0
  20. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/LICENSE +0 -0
  21. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/MANIFEST.in +0 -0
  22. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/README.md +0 -0
  23. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/benchmarks/README.md +0 -0
  24. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/__init__.py +0 -0
  25. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/core/__init__.py +0 -0
  26. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/core/base.py +0 -0
  27. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/core/base_coupling.py +0 -0
  28. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/core/base_equilibrium_field.py +0 -0
  29. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/core/base_equilibrium_interpolant.py +0 -0
  30. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/core/base_interpolant.py +0 -0
  31. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/core/base_loss.py +0 -0
  32. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/core/base_solver.py +0 -0
  33. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/core/base_velocity_field.py +0 -0
  34. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/datasets/__init__.py +0 -0
  35. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/datasets/radiograph.py +0 -0
  36. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/interpolants/__init__.py +0 -0
  37. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/interpolants/base.py +0 -0
  38. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/interpolants/equilibrium.py +0 -0
  39. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/interpolants/linear.py +0 -0
  40. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/interpolants/ot.py +0 -0
  41. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/interpolants/schrodinger_bridge.py +0 -0
  42. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/interpolants/variance_preserving.py +0 -0
  43. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/inverse/__init__.py +0 -0
  44. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/inverse/likelihood.py +0 -0
  45. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/inverse/operators.py +0 -0
  46. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/inverse/tweedie.py +0 -0
  47. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/losses/__init__.py +0 -0
  48. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/losses/conditional_flow_matching.py +0 -0
  49. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/losses/delta_alignment.py +0 -0
  50. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/losses/equilibrium_matching.py +0 -0
  51. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/losses/flow_matching.py +0 -0
  52. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/models/backbone.py +0 -0
  53. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/models/ema.py +0 -0
  54. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/models/projector.py +0 -0
  55. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/samplers/__init__.py +0 -0
  56. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/samplers/euler.py +0 -0
  57. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/solvers/__init__.py +0 -0
  58. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/solvers/euler.py +0 -0
  59. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/solvers/gradient_descent.py +0 -0
  60. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/solvers/heun.py +0 -0
  61. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/solvers/posterior_solver.py +0 -0
  62. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/trainer/__init__.py +0 -0
  63. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/trainer/coupling.py +0 -0
  64. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/trainer/data.py +0 -0
  65. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/trainer/loop.py +0 -0
  66. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/utils/__init__.py +0 -0
  67. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/utils/numerical.py +0 -0
  68. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/deltaflow/utils/ot.py +0 -0
  69. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/api/core.md +0 -0
  70. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/api/datasets.md +0 -0
  71. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/api/index.md +0 -0
  72. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/api/interpolants.md +0 -0
  73. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/api/inverse.md +0 -0
  74. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/api/losses.md +0 -0
  75. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/api/models.md +0 -0
  76. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/api/samplers.md +0 -0
  77. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/api/solvers.md +0 -0
  78. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/api/trainer.md +0 -0
  79. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/api/utils.md +0 -0
  80. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/assets/algorithm-comparison/comparison.gif +0 -0
  81. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/assets/algorithm-comparison/comparison.png +0 -0
  82. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/assets/algorithm-comparison/trajectories_comparison.png +0 -0
  83. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/assets/css/extra.css +0 -0
  84. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/assets/equilibrium-matching/energy_landscape.png +0 -0
  85. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/assets/equilibrium-matching/eqm_sampling.gif +0 -0
  86. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/assets/equilibrium-matching/gd_snapshots.png +0 -0
  87. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/assets/equilibrium-matching/gd_trajectories.png +0 -0
  88. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/assets/favicon.ico +0 -0
  89. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/assets/favicon.png +0 -0
  90. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/assets/guidance-alignment/features.png +0 -0
  91. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/assets/inverse-posterior/inverse_posterior.gif +0 -0
  92. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/assets/inverse-posterior/inverse_posterior.png +0 -0
  93. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/assets/js/mathjax.js +0 -0
  94. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/assets/landmark-detection/landmark_detection.png +0 -0
  95. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/assets/logo.svg +0 -0
  96. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/assets/minibatch-ot/minibatch_ot.gif +0 -0
  97. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/assets/minibatch-ot/minibatch_ot.png +0 -0
  98. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/assets/sampling-flow/flow.gif +0 -0
  99. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/assets/sampling-flow/snapshots.png +0 -0
  100. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/assets/sampling-flow/trajectories.png +0 -0
  101. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/assets/sampling-flow/velocity_field.png +0 -0
  102. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/assets/schrodinger-bridge/bridge_paths.gif +0 -0
  103. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/assets/schrodinger-bridge/bridge_paths.png +0 -0
  104. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/assets/schrodinger-bridge/sb_flow.gif +0 -0
  105. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/assets/schrodinger-bridge/sb_snapshots.png +0 -0
  106. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/assets/schrodinger-bridge/sb_trajectories.png +0 -0
  107. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/concepts/architecture.md +0 -0
  108. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/contributing.md +0 -0
  109. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/getting-started/installation.md +0 -0
  110. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/getting-started/quickstart.md +0 -0
  111. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/guides/delta-alignment.md +0 -0
  112. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/guides/flow-matching.md +0 -0
  113. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/guides/inverse-problems.md +0 -0
  114. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/guides/training.md +0 -0
  115. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/docs/index.md +0 -0
  116. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/examples/00-foundations/01-linear-interpolant/main.py +0 -0
  117. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/examples/10-sampling/01-euler-flow/main.py +0 -0
  118. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/examples/10-sampling/02-equilibrium-matching/main.py +0 -0
  119. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/examples/20-training/01-flow-matching/main.py +0 -0
  120. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/examples/20-training/02-delta-alignment/main.py +0 -0
  121. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/examples/30-inverse/01-posterior/main.py +0 -0
  122. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/examples/90-showcase/01-landmark-viz/main.py +0 -0
  123. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/examples/90-showcase/02-sampling-flow-viz/main.py +0 -0
  124. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/examples/90-showcase/03-minibatch-ot-viz/main.py +0 -0
  125. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/examples/90-showcase/04-inverse-posterior-viz/main.py +0 -0
  126. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/examples/90-showcase/05-schrodinger-bridge-viz/main.py +0 -0
  127. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/examples/90-showcase/06-algorithm-comparison/main.py +0 -0
  128. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/examples/90-showcase/07-landmark-detection/main.py +0 -0
  129. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/examples/90-showcase/08-guidance-alignment-pretraining/main.py +0 -0
  130. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/examples/90-showcase/09-equilibrium-matching-viz/main.py +0 -0
  131. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/examples/90-showcase/README.md +0 -0
  132. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/mkdocs.yml +0 -0
  133. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/overrides/main.html +0 -0
  134. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/pyproject.toml +0 -0
  135. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/requirements-docs.txt +0 -0
  136. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/requirements.txt +0 -0
  137. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/setup.cfg +0 -0
  138. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/tests/__init__.py +0 -0
  139. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/tests/conftest.py +0 -0
  140. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/tests/test_coupling.py +0 -0
  141. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/tests/test_equilibrium.py +0 -0
  142. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/tests/test_interpolants.py +0 -0
  143. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/tests/test_interpolants_new.py +0 -0
  144. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/tests/test_inverse.py +0 -0
  145. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/tests/test_losses.py +0 -0
  146. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/tests/test_models.py +0 -0
  147. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/tests/test_samplers.py +0 -0
  148. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/tests/test_solvers.py +0 -0
  149. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/tests/test_trainer.py +0 -0
  150. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/torchdeltaflow.egg-info/dependency_links.txt +0 -0
  151. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/torchdeltaflow.egg-info/requires.txt +0 -0
  152. {torchdeltaflow-0.2.3.dev21 → torchdeltaflow-0.2.3.dev23}/torchdeltaflow.egg-info/top_level.txt +0 -0
@@ -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.3.dev23
4
4
  Summary: Flow matching, mini-batch optimal transport coupling, and posterior sampling for inverse problems in PyTorch
5
5
  Author-email: Phrugsa Limbunlom <phrugsa.lim@gmail.com>
6
6
  Maintainer-email: Phrugsa Limbunlom <phrugsa.lim@gmail.com>
@@ -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
+ ]
@@ -146,6 +146,42 @@ difference feature.
146
146
 
147
147
  ---
148
148
 
149
+ ## 20 training, conditional generation with a DiT
150
+
151
+ [`examples/20-training/03-conditional-dit/main.py`](https://github.com/phrugsa-limbunlom/deltaflow/blob/main/examples/20-training/03-conditional-dit/main.py)
152
+
153
+ The native transformer path. A small `DiT` velocity field conditions on a class
154
+ label through **adaLN-Zero** (Peebles & Xie, 2023), the shared ``time + class``
155
+ embedding drives per-block ``(shift, scale, gate)`` modulation, and the gates
156
+ start at zero so the field begins as the identity. Two toy classes (bright-top
157
+ versus bright-bottom 8x8 images) are learned with conditional flow matching, then
158
+ sampled per class with classifier-free guidance.
159
+
160
+ ```python
161
+ from deltaflow.models import DiT
162
+
163
+ field = DiT(input_size=8, patch_size=2, in_channels=1, hidden_size=64,
164
+ depth=2, num_heads=4, num_classes=2, class_dropout_prob=0.1)
165
+
166
+ loss = FlowMatchingLoss()(field, x1, y=y) # label rides the opaque **cond channel
167
+
168
+ # guided sampling: extrapolate between the conditional and null-token velocity
169
+ v = field.forward_with_cfg(x, t, y=torch.zeros(64, dtype=torch.long), cfg_scale=2.0)
170
+ ```
171
+
172
+ What to notice.
173
+
174
+ - The label enters through the same opaque ``**cond`` channel every DeltaFlow
175
+ loss and solver forwards unchanged, here as ``y``. No loss or solver needs to
176
+ know about classes.
177
+ - `LabelEmbedding` carries a learned **null** token. Train-time label dropout
178
+ teaches one network both the conditional and unconditional velocity, which is
179
+ exactly what `DiT.forward_with_cfg` extrapolates between at sample time.
180
+ - The printed per-class pixel means separate cleanly (class 0 bright on top,
181
+ class 1 bright on bottom), confirming the conditioning took effect.
182
+
183
+ ---
184
+
149
185
  ## 30 inverse, posterior sampling
150
186
 
151
187
  [`examples/30-inverse/01-posterior/main.py`](https://github.com/phrugsa-limbunlom/deltaflow/blob/main/examples/30-inverse/01-posterior/main.py)
@@ -0,0 +1,72 @@
1
+ """
2
+ 20-training/03-conditional-dit: class-conditional generation with a DiT
3
+ velocity field and adaLN-Zero conditioning (issue #11).
4
+
5
+ Two classes draw tiny 8x8 "images": class 0 is a bright top half, class 1 is a
6
+ bright bottom half. A small `DiT` is trained with conditional flow matching,
7
+ then sampled per class with classifier-free guidance. The per-class pixel means
8
+ confirm the label conditioning took effect.
9
+
10
+ Run: python examples/20-training/03-conditional-dit/main.py
11
+ """
12
+
13
+ import torch
14
+
15
+ from deltaflow.losses import FlowMatchingLoss
16
+ from deltaflow.models import DiT
17
+ from deltaflow.solvers import EulerSolver
18
+
19
+ IMG = 8
20
+ N_CLASSES = 2
21
+
22
+
23
+ def sample_batch(n: int):
24
+ """Return (images, labels): class 0 bright on top, class 1 bright on bottom."""
25
+ y = torch.randint(0, N_CLASSES, (n,))
26
+ x = torch.zeros(n, 1, IMG, IMG)
27
+ half = IMG // 2
28
+ top = y == 0
29
+ x[top, :, :half, :] = 1.0
30
+ x[~top, :, half:, :] = 1.0
31
+ return x + 0.05 * torch.randn_like(x), y
32
+
33
+
34
+ def main():
35
+ torch.manual_seed(0)
36
+ field = DiT(
37
+ input_size=IMG,
38
+ patch_size=2,
39
+ in_channels=1,
40
+ hidden_size=64,
41
+ depth=2,
42
+ num_heads=4,
43
+ num_classes=N_CLASSES,
44
+ class_dropout_prob=0.1,
45
+ )
46
+ loss_fn = FlowMatchingLoss()
47
+ opt = torch.optim.Adam(field.parameters(), lr=2e-3)
48
+
49
+ field.train()
50
+ for step in range(300):
51
+ x1, y = sample_batch(64)
52
+ loss = loss_fn(field, x1, y=y)
53
+ opt.zero_grad()
54
+ loss.backward()
55
+ opt.step()
56
+ if step % 100 == 0:
57
+ print(f"step={step:4d} loss={loss.item():.4f}")
58
+
59
+ # Classifier-free-guided sampling, conditioned on each class.
60
+ field.eval()
61
+ for cls in range(N_CLASSES):
62
+ y = torch.full((64,), cls)
63
+ solver = EulerSolver(lambda x, t, y=y: field.forward_with_cfg(x, t, y, cfg_scale=2.0))
64
+ samples = solver.sample(torch.randn(64, 1, IMG, IMG), n_steps=50, show_progress=False)
65
+ half = IMG // 2
66
+ top_mean = samples[:, :, :half, :].mean().item()
67
+ bot_mean = samples[:, :, half:, :].mean().item()
68
+ print(f"class {cls}: top_mean={top_mean:+.3f} bottom_mean={bot_mean:+.3f}")
69
+
70
+
71
+ if __name__ == "__main__":
72
+ main()