diffusers-workflow 0.4.0__py3-none-any.whl

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 (260) hide show
  1. diffusers_workflow-0.4.0.dist-info/METADATA +318 -0
  2. diffusers_workflow-0.4.0.dist-info/RECORD +260 -0
  3. diffusers_workflow-0.4.0.dist-info/WHEEL +5 -0
  4. diffusers_workflow-0.4.0.dist-info/entry_points.txt +7 -0
  5. diffusers_workflow-0.4.0.dist-info/licenses/LICENSE +201 -0
  6. diffusers_workflow-0.4.0.dist-info/top_level.txt +2 -0
  7. dw/__init__.py +440 -0
  8. dw/adapter_compatibility.py +226 -0
  9. dw/arguments.py +1231 -0
  10. dw/assessment_rules.py +159 -0
  11. dw/assets.py +130 -0
  12. dw/cache_blocks.json +16 -0
  13. dw/cache_blocks.py +146 -0
  14. dw/community_pipelines/pipeline_flux_rf_inversion.py +1184 -0
  15. dw/content_types.py +150 -0
  16. dw/dissolve_frame_errors.py +121 -0
  17. dw/docs/ACCELERATION.md +352 -0
  18. dw/docs/AGENT_LOOP.md +95 -0
  19. dw/docs/DEPENDENCIES.md +91 -0
  20. dw/docs/IP_ADAPTER.md +109 -0
  21. dw/docs/LORAS.md +131 -0
  22. dw/docs/MCP.md +517 -0
  23. dw/docs/PROMPT_WEIGHTING.md +78 -0
  24. dw/docs/QUANTIZATION.md +230 -0
  25. dw/docs/RECIPES_24GB.md +201 -0
  26. dw/docs/RELEASING.md +195 -0
  27. dw/docs/REMOTE.md +140 -0
  28. dw/docs/REPL_COMMANDS.md +121 -0
  29. dw/docs/REPL_WORKER_GUIDE.md +51 -0
  30. dw/docs/SECURITY.md +272 -0
  31. dw/docs/SECURITY_QUICKREF.md +112 -0
  32. dw/docs/SERVER.md +679 -0
  33. dw/docs/TASKS.md +1741 -0
  34. dw/docs/TESTING.md +71 -0
  35. dw/docs/WORKFLOW_GUIDE.md +2038 -0
  36. dw/docs/WORKSPACES.md +316 -0
  37. dw/download_watch.py +335 -0
  38. dw/elision.py +306 -0
  39. dw/events.py +275 -0
  40. dw/for_each.py +409 -0
  41. dw/host_memory.py +258 -0
  42. dw/host_memory_projection.py +230 -0
  43. dw/hub_cache.py +432 -0
  44. dw/introspection.py +1228 -0
  45. dw/kernel_availability.py +208 -0
  46. dw/locations.py +599 -0
  47. dw/log_setup.py +45 -0
  48. dw/loudness.py +82 -0
  49. dw/media_audio.py +217 -0
  50. dw/media_frames.py +367 -0
  51. dw/media_info.py +297 -0
  52. dw/pipeline_processors/chain.py +821 -0
  53. dw/pipeline_processors/config_objects.py +237 -0
  54. dw/pipeline_processors/pipeline.py +2297 -0
  55. dw/pipeline_processors/remote.py +46 -0
  56. dw/plan.py +920 -0
  57. dw/previous_results.py +411 -0
  58. dw/probe_paths.py +59 -0
  59. dw/prompt_schema.json +48 -0
  60. dw/prompt_weighting.py +378 -0
  61. dw/prompts.py +159 -0
  62. dw/realize.py +250 -0
  63. dw/reference_limits.py +215 -0
  64. dw/reference_names.py +125 -0
  65. dw/repl.py +338 -0
  66. dw/repl_commands.py +836 -0
  67. dw/repl_worker.py +159 -0
  68. dw/result.py +1720 -0
  69. dw/result_fps.py +82 -0
  70. dw/run.py +162 -0
  71. dw/runs.py +768 -0
  72. dw/scalar_result_validation.py +97 -0
  73. dw/schema.py +283 -0
  74. dw/security.py +1038 -0
  75. dw/select_validation.py +115 -0
  76. dw/serve.py +277 -0
  77. dw/server/__init__.py +2 -0
  78. dw/server/app.py +4586 -0
  79. dw/server/assess.py +132 -0
  80. dw/server/catalog_shape.py +487 -0
  81. dw/server/enhancers.py +129 -0
  82. dw/server/exports.py +480 -0
  83. dw/server/guides.py +257 -0
  84. dw/server/jobs.py +1561 -0
  85. dw/server/mcp_mount.py +95 -0
  86. dw/server/netinfo.py +124 -0
  87. dw/server/observed_cost.py +379 -0
  88. dw/server/sysinfo.py +71 -0
  89. dw/server/ui/assets/abap-08VXUWAP.js +1 -0
  90. dw/server/ui/assets/apex-BWPQTe0t.js +1 -0
  91. dw/server/ui/assets/azcli-Bc_sGQ0U.js +1 -0
  92. dw/server/ui/assets/bat-i0X4ZdIN.js +1 -0
  93. dw/server/ui/assets/bicep-B5-_aFwp.js +2 -0
  94. dw/server/ui/assets/cameligo-DMUM7wLl.js +1 -0
  95. dw/server/ui/assets/clojure-Cm7r79vr.js +1 -0
  96. dw/server/ui/assets/codicon-Brq4_Ui5.ttf +0 -0
  97. dw/server/ui/assets/coffee-Ba7i2nA0.js +1 -0
  98. dw/server/ui/assets/cpp-C7h46wYY.js +1 -0
  99. dw/server/ui/assets/csharp-BKxtCVv1.js +1 -0
  100. dw/server/ui/assets/csp-bTuwJoIa.js +1 -0
  101. dw/server/ui/assets/css-DIMkf-bt.js +3 -0
  102. dw/server/ui/assets/css.worker-B3ciXF_0.js +93 -0
  103. dw/server/ui/assets/cssMode-CPznxfY8.js +1 -0
  104. dw/server/ui/assets/cypher-CVaqCwHa.js +1 -0
  105. dw/server/ui/assets/dart-onAF5SnQ.js +1 -0
  106. dw/server/ui/assets/dockerfile-DZFCIeNp.js +1 -0
  107. dw/server/ui/assets/ecl-D05T4iGw.js +1 -0
  108. dw/server/ui/assets/editor-jjEx9u7D.css +1 -0
  109. dw/server/ui/assets/editor.api-CpWcotrd.js +847 -0
  110. dw/server/ui/assets/editor.worker-q-txB4vs.js +30 -0
  111. dw/server/ui/assets/elixir-6RTg0lbw.js +1 -0
  112. dw/server/ui/assets/flow9-C5_-GSwl.js +1 -0
  113. dw/server/ui/assets/freemarker2-CXtRM8N4.js +3 -0
  114. dw/server/ui/assets/fsharp-C8Ef5oNN.js +1 -0
  115. dw/server/ui/assets/go-C-y9NEjX.js +1 -0
  116. dw/server/ui/assets/graphql-fmXr3nnJ.js +1 -0
  117. dw/server/ui/assets/handlebars-N7x-6NMY.js +1 -0
  118. dw/server/ui/assets/hcl-CpzslTdj.js +1 -0
  119. dw/server/ui/assets/html-PhsdjHSr.js +1 -0
  120. dw/server/ui/assets/html.worker-C93Ht9o9.js +506 -0
  121. dw/server/ui/assets/htmlMode-Dgj0SEok.js +1 -0
  122. dw/server/ui/assets/index-3Vw6WAPW.css +1 -0
  123. dw/server/ui/assets/index-DgrYhQd9.js +43 -0
  124. dw/server/ui/assets/ini-sBoK_t0W.js +1 -0
  125. dw/server/ui/assets/java-BEtHBSE6.js +1 -0
  126. dw/server/ui/assets/javascript-BJqN9Qhv.js +1 -0
  127. dw/server/ui/assets/json.worker-B2V3pomh.js +62 -0
  128. dw/server/ui/assets/jsonMode-DbM4SWSv.js +7 -0
  129. dw/server/ui/assets/julia-Bri6UV-V.js +1 -0
  130. dw/server/ui/assets/kotlin-BOotOW0E.js +1 -0
  131. dw/server/ui/assets/less-B9JPFI3C.js +2 -0
  132. dw/server/ui/assets/lexon-CfSJPG6W.js +1 -0
  133. dw/server/ui/assets/liquid-BWr8lEc4.js +1 -0
  134. dw/server/ui/assets/lspLanguageFeatures-C1iGuDyZ.js +4 -0
  135. dw/server/ui/assets/lua-CsQS60Ue.js +1 -0
  136. dw/server/ui/assets/m3-D-oSqn_W.js +1 -0
  137. dw/server/ui/assets/markdown-Cimd5fb3.js +1 -0
  138. dw/server/ui/assets/mdx-DAdMi_0p.js +1 -0
  139. dw/server/ui/assets/mips-CIPQ_RoX.js +1 -0
  140. dw/server/ui/assets/monaco--ixms01u.css +1 -0
  141. dw/server/ui/assets/monaco-BGCeEqaw.js +56 -0
  142. dw/server/ui/assets/msdax-DauUninz.js +1 -0
  143. dw/server/ui/assets/mysql-SOo6toE5.js +1 -0
  144. dw/server/ui/assets/objective-c-FvmIjYaQ.js +1 -0
  145. dw/server/ui/assets/pascal-DrH0SRf2.js +1 -0
  146. dw/server/ui/assets/pascaligo-D-ptJ9y-.js +1 -0
  147. dw/server/ui/assets/perl-oz_6vUea.js +1 -0
  148. dw/server/ui/assets/pgsql-DTj74zXo.js +1 -0
  149. dw/server/ui/assets/php-nr791fC2.js +1 -0
  150. dw/server/ui/assets/pla-CopQ2nXW.js +1 -0
  151. dw/server/ui/assets/postiats-43DmfD33.js +1 -0
  152. dw/server/ui/assets/powerquery-D3hlyOfw.js +1 -0
  153. dw/server/ui/assets/powershell-DmHpPYUd.js +1 -0
  154. dw/server/ui/assets/protobuf-C531GsRP.js +2 -0
  155. dw/server/ui/assets/pug-Z5eAx3Zn.js +1 -0
  156. dw/server/ui/assets/python-Bcn70HdC.js +1 -0
  157. dw/server/ui/assets/qsharp-DkqhCAOL.js +1 -0
  158. dw/server/ui/assets/r-BwWrilGY.js +1 -0
  159. dw/server/ui/assets/razor-D1HmNnby.js +1 -0
  160. dw/server/ui/assets/redis-ClamHrr6.js +1 -0
  161. dw/server/ui/assets/redshift-DT7zqm-g.js +1 -0
  162. dw/server/ui/assets/restructuredtext-BYgofb2h.js +1 -0
  163. dw/server/ui/assets/ruby-DezsRK8O.js +1 -0
  164. dw/server/ui/assets/rust-DdL9SqIa.js +1 -0
  165. dw/server/ui/assets/sb-CcwsVR0C.js +1 -0
  166. dw/server/ui/assets/scala-DHpiXF5c.js +1 -0
  167. dw/server/ui/assets/scheme-BeGwcela.js +1 -0
  168. dw/server/ui/assets/scss-gp-XZpBa.js +3 -0
  169. dw/server/ui/assets/shell-CC2rA5mh.js +1 -0
  170. dw/server/ui/assets/solidity-BEEn4gHE.js +1 -0
  171. dw/server/ui/assets/sophia-CRfGWb83.js +1 -0
  172. dw/server/ui/assets/sparql-D_Lu-MrJ.js +1 -0
  173. dw/server/ui/assets/sql-NEE52Syq.js +1 -0
  174. dw/server/ui/assets/st-DbInun42.js +1 -0
  175. dw/server/ui/assets/swift-Bxkupp3x.js +1 -0
  176. dw/server/ui/assets/systemverilog-Bz4Y3fRF.js +1 -0
  177. dw/server/ui/assets/tcl-DISqw1ZD.js +1 -0
  178. dw/server/ui/assets/ts.worker-D7T1-Ig5.js +67738 -0
  179. dw/server/ui/assets/tsMode-D6u0XmOW.js +11 -0
  180. dw/server/ui/assets/twig-De2hgUGE.js +1 -0
  181. dw/server/ui/assets/typescript-BU6v-LMV.js +1 -0
  182. dw/server/ui/assets/typespec-B8J7ngcE.js +1 -0
  183. dw/server/ui/assets/vb-DV3o63ZY.js +1 -0
  184. dw/server/ui/assets/wgsl-DpFanUEy.js +298 -0
  185. dw/server/ui/assets/workers-Cn7cTUKr.js +1 -0
  186. dw/server/ui/assets/xml--0LP2Lwk.js +1 -0
  187. dw/server/ui/assets/yaml-mpBg9jnt.js +1 -0
  188. dw/server/ui/index.html +17 -0
  189. dw/server/updater.py +192 -0
  190. dw/settings.py +98 -0
  191. dw/shot_span_preflight.py +116 -0
  192. dw/shots.py +359 -0
  193. dw/slice_preflight.py +148 -0
  194. dw/step.py +187 -0
  195. dw/step_cache.py +442 -0
  196. dw/subfolders.py +107 -0
  197. dw/task_domains.py +307 -0
  198. dw/tasks/assess.py +826 -0
  199. dw/tasks/audio_transcription.py +88 -0
  200. dw/tasks/audio_utils.py +1862 -0
  201. dw/tasks/background_remover.py +43 -0
  202. dw/tasks/borders.py +113 -0
  203. dw/tasks/compose_text.py +74 -0
  204. dw/tasks/concat_videos.py +300 -0
  205. dw/tasks/depth_estimator.py +54 -0
  206. dw/tasks/diffusion_upscale.py +109 -0
  207. dw/tasks/dissolve_videos.py +342 -0
  208. dw/tasks/format_messages.py +24 -0
  209. dw/tasks/gather.py +173 -0
  210. dw/tasks/grade.py +97 -0
  211. dw/tasks/image_to_text.py +43 -0
  212. dw/tasks/image_utils.py +764 -0
  213. dw/tasks/interpolate_frames.py +252 -0
  214. dw/tasks/judge.py +68 -0
  215. dw/tasks/model_cache.py +55 -0
  216. dw/tasks/pair_audio.py +268 -0
  217. dw/tasks/qr_code.py +19 -0
  218. dw/tasks/restore_faces.py +175 -0
  219. dw/tasks/rife_model.py +192 -0
  220. dw/tasks/segment.py +121 -0
  221. dw/tasks/select.py +111 -0
  222. dw/tasks/speech_generation.py +228 -0
  223. dw/tasks/stabilize.py +129 -0
  224. dw/tasks/task.py +920 -0
  225. dw/tasks/tensor_image.py +57 -0
  226. dw/tasks/text_generation.py +169 -0
  227. dw/tasks/text_sections.py +80 -0
  228. dw/tasks/upscale.py +203 -0
  229. dw/tasks/video_utils.py +624 -0
  230. dw/tasks/zoe_depth.py +71 -0
  231. dw/teacache.py +381 -0
  232. dw/teacache_models.json +99 -0
  233. dw/test.py +29 -0
  234. dw/type_helpers.py +231 -0
  235. dw/validate.py +68 -0
  236. dw/variable_constraints.py +444 -0
  237. dw/variables.py +443 -0
  238. dw/video_extensions.py +141 -0
  239. dw/vram_estimate.py +116 -0
  240. dw/worker.py +764 -0
  241. dw/workflow.py +2007 -0
  242. dw/workflow_schema.json +1346 -0
  243. dw/workflow_sources.py +383 -0
  244. dw/workflows/h3_context_ir.json +57 -0
  245. dw/workflows/test.json +31 -0
  246. dw/workspace.py +730 -0
  247. dw_mcp/__init__.py +6 -0
  248. dw_mcp/__main__.py +133 -0
  249. dw_mcp/assets.py +336 -0
  250. dw_mcp/authoring.py +114 -0
  251. dw_mcp/catalog.py +360 -0
  252. dw_mcp/client.py +486 -0
  253. dw_mcp/diagnose.py +371 -0
  254. dw_mcp/exports.py +84 -0
  255. dw_mcp/guides.py +35 -0
  256. dw_mcp/media.py +638 -0
  257. dw_mcp/models.py +97 -0
  258. dw_mcp/prompts.py +104 -0
  259. dw_mcp/server.py +1343 -0
  260. dw_mcp/workspaces.py +212 -0
dw/teacache.py ADDED
@@ -0,0 +1,381 @@
1
+ """
2
+ TeaCache - Training-free inference acceleration for diffusion transformers.
3
+
4
+ Caches intermediate transformer computations and skips redundant steps
5
+ when the input hasn't changed significantly between timesteps.
6
+
7
+ Based on: https://github.com/ali-vilab/TeaCache
8
+ Adapted from: https://github.com/Teriks/dgenerate (Apache 2.0)
9
+
10
+ Implemented: Flux (FluxTransformer2DModel)
11
+ Registry includes: Mochi, LTX-Video, CogVideoX, Lumina2, HunyuanVideo, Wan2.1
12
+ (these require custom forward functions to be added)
13
+
14
+ Each model requires a custom forward function because transformer architectures
15
+ differ. The core caching algorithm is the same: extract a signal from the first
16
+ block's normalization, compare via polynomial rescaling, skip if below threshold.
17
+ """
18
+
19
+ import json
20
+ import typing
21
+ import logging
22
+ from pathlib import Path
23
+ from contextlib import contextmanager
24
+
25
+ import torch
26
+ import numpy as np
27
+ from diffusers.models.modeling_outputs import Transformer2DModelOutput
28
+ from diffusers.utils import (
29
+ USE_PEFT_BACKEND,
30
+ scale_lora_layers,
31
+ unscale_lora_layers,
32
+ )
33
+
34
+ logger = logging.getLogger("dw")
35
+
36
+
37
+ # ---------------------------------------------------------------------------
38
+ # Model registry loaded from JSON
39
+ # ---------------------------------------------------------------------------
40
+
41
+ _REGISTRY_PATH = Path(__file__).parent / "teacache_models.json"
42
+
43
+
44
+ def _load_registry():
45
+ """Load the model registry from the JSON file."""
46
+ with open(_REGISTRY_PATH) as f:
47
+ return json.load(f)
48
+
49
+
50
+ def _get_model_info(transformer, variant=None):
51
+ """Look up model info from registry.
52
+
53
+ Args:
54
+ transformer: The transformer model instance.
55
+ variant: Optional explicit variant name (e.g., "wan2.1_t2v_1.3b").
56
+ If None, uses class_defaults mapping.
57
+
58
+ Returns:
59
+ dict with coefficients, default_threshold, threshold_guide.
60
+ """
61
+ registry = _load_registry()
62
+ class_name = transformer.__class__.__name__
63
+
64
+ if variant is not None:
65
+ info = registry["models"].get(variant)
66
+ if info is None:
67
+ available = ", ".join(registry["models"].keys())
68
+ raise ValueError(
69
+ f"TeaCache variant '{variant}' not found. Available: {available}"
70
+ )
71
+ return info
72
+
73
+ # Look up default variant for this class
74
+ default_variant = registry["class_defaults"].get(class_name)
75
+ if default_variant is None:
76
+ supported_classes = ", ".join(registry["class_defaults"].keys())
77
+ raise ValueError(
78
+ f"TeaCache does not support {class_name}. Supported: {supported_classes}"
79
+ )
80
+
81
+ return registry["models"][default_variant]
82
+
83
+
84
+ # ---------------------------------------------------------------------------
85
+ # Forward function factories, one per supported transformer architecture.
86
+ # ---------------------------------------------------------------------------
87
+
88
+
89
+ def _create_flux_teacache_forward(num_inference_steps, rel_l1_thresh, coefficients):
90
+ """Create TeaCache forward for FluxTransformer2DModel."""
91
+ cnt = 0
92
+ accumulated_rel_l1_distance = 0
93
+ previous_modulated_input = None
94
+ previous_residual = None
95
+ previous_timestep = None
96
+ rescale_func = np.poly1d(coefficients)
97
+
98
+ def teacache_forward(
99
+ self,
100
+ hidden_states: torch.Tensor,
101
+ encoder_hidden_states: torch.Tensor = None,
102
+ pooled_projections: torch.Tensor = None,
103
+ timestep: torch.LongTensor = None,
104
+ img_ids: torch.Tensor = None,
105
+ txt_ids: torch.Tensor = None,
106
+ guidance: torch.Tensor = None,
107
+ joint_attention_kwargs: typing.Optional[typing.Dict[str, typing.Any]] = None,
108
+ controlnet_block_samples=None,
109
+ controlnet_single_block_samples=None,
110
+ return_dict: bool = True,
111
+ controlnet_blocks_repeat: bool = False,
112
+ ) -> typing.Union[torch.FloatTensor, Transformer2DModelOutput]:
113
+ nonlocal \
114
+ cnt, \
115
+ accumulated_rel_l1_distance, \
116
+ previous_modulated_input, \
117
+ previous_residual, \
118
+ previous_timestep
119
+
120
+ # TeaCache assumes exactly one transformer forward call per denoising
121
+ # step. Pipelines running true classifier-free guidance (e.g. Flux with
122
+ # negative_prompt + true_cfg_scale > 1) call the transformer twice per
123
+ # step -- once for the conditional pass and once for the unconditional
124
+ # pass -- using the identical timestep both times. That second call
125
+ # would silently share/corrupt previous_modulated_input and
126
+ # previous_residual across the two passes, so detect it and fail loudly
127
+ # instead of producing a corrupted image.
128
+ if (
129
+ timestep is not None
130
+ and previous_timestep is not None
131
+ and timestep.shape == previous_timestep.shape
132
+ and torch.equal(timestep, previous_timestep)
133
+ ):
134
+ raise RuntimeError(
135
+ "TeaCache does not support true classifier-free guidance "
136
+ "(negative_prompt with true_cfg_scale > 1); disable one of them. "
137
+ "Detected two transformer forward calls with an identical "
138
+ "timestep within a single denoising step, which would corrupt "
139
+ "TeaCache's cached state."
140
+ )
141
+ if timestep is not None:
142
+ previous_timestep = timestep.detach().clone()
143
+
144
+ if joint_attention_kwargs is not None:
145
+ joint_attention_kwargs = joint_attention_kwargs.copy()
146
+ lora_scale = joint_attention_kwargs.pop("scale", 1.0)
147
+ else:
148
+ lora_scale = 1.0
149
+
150
+ if USE_PEFT_BACKEND:
151
+ scale_lora_layers(self, lora_scale)
152
+
153
+ hidden_states = self.x_embedder(hidden_states)
154
+
155
+ timestep = timestep.to(hidden_states.dtype) * 1000
156
+ if guidance is not None:
157
+ guidance = guidance.to(hidden_states.dtype) * 1000
158
+
159
+ temb = (
160
+ self.time_text_embed(timestep, pooled_projections)
161
+ if guidance is None
162
+ else self.time_text_embed(timestep, guidance, pooled_projections)
163
+ )
164
+ encoder_hidden_states = self.context_embedder(encoder_hidden_states)
165
+
166
+ if txt_ids.ndim == 3:
167
+ txt_ids = txt_ids[0]
168
+ if img_ids.ndim == 3:
169
+ img_ids = img_ids[0]
170
+
171
+ ids = torch.cat((txt_ids, img_ids), dim=0)
172
+ image_rotary_emb = self.pos_embed(ids)
173
+
174
+ if (
175
+ joint_attention_kwargs is not None
176
+ and "ip_adapter_image_embeds" in joint_attention_kwargs
177
+ ):
178
+ ip_adapter_image_embeds = joint_attention_kwargs.pop(
179
+ "ip_adapter_image_embeds"
180
+ )
181
+ ip_hidden_states = self.encoder_hid_proj(ip_adapter_image_embeds)
182
+ joint_attention_kwargs.update({"ip_hidden_states": ip_hidden_states})
183
+
184
+ # TeaCache: extract cache signal from first block's normalization
185
+ inp = hidden_states.clone()
186
+ temb_ = temb.clone()
187
+ modulated_inp, gate_msa, shift_mlp, scale_mlp, gate_mlp = (
188
+ self.transformer_blocks[0].norm1(inp, emb=temb_)
189
+ )
190
+
191
+ # Decide whether to compute or reuse cached result
192
+ if cnt == 0 or cnt == num_inference_steps - 1:
193
+ should_calc = True
194
+ accumulated_rel_l1_distance = 0
195
+ else:
196
+ relative_diff = (
197
+ (
198
+ (modulated_inp - previous_modulated_input).abs().mean()
199
+ / previous_modulated_input.abs().mean()
200
+ )
201
+ .cpu()
202
+ .item()
203
+ )
204
+ accumulated_rel_l1_distance += rescale_func(relative_diff)
205
+
206
+ if accumulated_rel_l1_distance < rel_l1_thresh:
207
+ should_calc = False
208
+ else:
209
+ should_calc = True
210
+ accumulated_rel_l1_distance = 0
211
+
212
+ previous_modulated_input = modulated_inp
213
+ cnt += 1
214
+ if cnt == num_inference_steps:
215
+ cnt = 0
216
+
217
+ if not should_calc:
218
+ hidden_states += previous_residual
219
+ else:
220
+ ori_hidden_states = hidden_states.clone()
221
+
222
+ # No gradient-checkpointing branch: this forward only runs under
223
+ # Pipeline.run's @torch.inference_mode(), so grads are never enabled
224
+ for index_block, block in enumerate(self.transformer_blocks):
225
+ encoder_hidden_states, hidden_states = block(
226
+ hidden_states=hidden_states,
227
+ encoder_hidden_states=encoder_hidden_states,
228
+ temb=temb,
229
+ image_rotary_emb=image_rotary_emb,
230
+ joint_attention_kwargs=joint_attention_kwargs,
231
+ )
232
+
233
+ if controlnet_block_samples is not None:
234
+ interval_control = len(self.transformer_blocks) / len(
235
+ controlnet_block_samples
236
+ )
237
+ interval_control = int(np.ceil(interval_control))
238
+ if controlnet_blocks_repeat:
239
+ hidden_states = (
240
+ hidden_states
241
+ + controlnet_block_samples[
242
+ index_block % len(controlnet_block_samples)
243
+ ]
244
+ )
245
+ else:
246
+ hidden_states = (
247
+ hidden_states
248
+ + controlnet_block_samples[index_block // interval_control]
249
+ )
250
+
251
+ for index_block, block in enumerate(self.single_transformer_blocks):
252
+ encoder_hidden_states, hidden_states = block(
253
+ hidden_states=hidden_states,
254
+ encoder_hidden_states=encoder_hidden_states,
255
+ temb=temb,
256
+ image_rotary_emb=image_rotary_emb,
257
+ joint_attention_kwargs=joint_attention_kwargs,
258
+ )
259
+
260
+ if controlnet_single_block_samples is not None:
261
+ interval_control = len(self.single_transformer_blocks) / len(
262
+ controlnet_single_block_samples
263
+ )
264
+ interval_control = int(np.ceil(interval_control))
265
+ hidden_states[:, encoder_hidden_states.shape[1] :, ...] = (
266
+ hidden_states[:, encoder_hidden_states.shape[1] :, ...]
267
+ + controlnet_single_block_samples[
268
+ index_block // interval_control
269
+ ]
270
+ )
271
+
272
+ previous_residual = hidden_states - ori_hidden_states
273
+
274
+ hidden_states = self.norm_out(hidden_states, temb)
275
+ output = self.proj_out(hidden_states)
276
+
277
+ if USE_PEFT_BACKEND:
278
+ unscale_lora_layers(self, lora_scale)
279
+
280
+ if not return_dict:
281
+ return (output,)
282
+
283
+ return Transformer2DModelOutput(sample=output)
284
+
285
+ return teacache_forward
286
+
287
+
288
+ # Map transformer class names to their forward factory functions.
289
+ # Models in the JSON registry without a factory here will get an informative error.
290
+ _FORWARD_FACTORIES = {
291
+ "FluxTransformer2DModel": _create_flux_teacache_forward,
292
+ }
293
+
294
+
295
+ # ---------------------------------------------------------------------------
296
+ # Public API
297
+ # ---------------------------------------------------------------------------
298
+
299
+
300
+ @contextmanager
301
+ def teacache_context(
302
+ pipeline, num_inference_steps, rel_l1_thresh=None, coefficients=None, variant=None
303
+ ):
304
+ """Context manager that enables TeaCache on a pipeline's transformer.
305
+
306
+ Auto-detects the transformer type and applies the appropriate
307
+ TeaCache forward function. Restores original forward on exit.
308
+
309
+ Args:
310
+ pipeline: A DiffusionPipeline with a .transformer attribute
311
+ num_inference_steps: Number of inference steps (must match pipeline call)
312
+ rel_l1_thresh: Cache threshold override. If None, uses model default.
313
+ Higher = more speedup, more quality loss.
314
+ coefficients: Polynomial coefficients override. If None, uses model default.
315
+ List of 5 floats for np.poly1d rescaling function.
316
+ variant: Explicit model variant name (e.g., "wan2.1_t2v_1.3b").
317
+ Required when a transformer class has multiple variants (CogVideoX, Wan).
318
+ If None, uses class_defaults from the registry.
319
+ """
320
+ transformer = pipeline.transformer
321
+ class_name = transformer.__class__.__name__
322
+
323
+ # Look up model info from registry
324
+ model_info = _get_model_info(transformer, variant)
325
+
326
+ # Check we have a forward implementation for this class
327
+ factory = _FORWARD_FACTORIES.get(class_name)
328
+ if factory is None:
329
+ supported = ", ".join(_FORWARD_FACTORIES.keys())
330
+ raise ValueError(
331
+ f"No TeaCache forward implementation for {class_name}. "
332
+ f"Implemented: {supported}. "
333
+ f"The model is in the registry but needs a custom forward function."
334
+ )
335
+
336
+ # Use overrides or defaults
337
+ if rel_l1_thresh is None:
338
+ rel_l1_thresh = model_info["default_threshold"]
339
+ if coefficients is None:
340
+ coefficients = model_info["coefficients"]
341
+
342
+ teacache_forward_fn = factory(num_inference_steps, rel_l1_thresh, coefficients)
343
+
344
+ # accelerate's enable_model_cpu_offload/enable_sequential_cpu_offload installs
345
+ # an AlignDevicesHook via add_hook_to_module (accelerate/hooks.py), which
346
+ # replaces transformer.forward with a wrapper closing over module and the
347
+ # true original forward (stashed as transformer._old_forward, still bound to
348
+ # the instance). That wrapper is what moves the module's weights to the
349
+ # execution device in its pre_forward before calling _old_forward. If we
350
+ # clobber transformer.forward like the no-hook path below, we remove that
351
+ # wrapper entirely: the CPU-resident module then receives CUDA inputs and
352
+ # raises a device-mismatch RuntimeError. Instead, when a hook is present we
353
+ # wrap what the hook considers "the real forward" -- _old_forward -- so the
354
+ # call chain stays hook.forward -> pre_forward (places weights) ->
355
+ # teacache_forward -> post_forward.
356
+ has_hook = hasattr(transformer, "_hf_hook") and hasattr(transformer, "_old_forward")
357
+
358
+ if has_hook:
359
+ original_forward = transformer._old_forward
360
+ transformer._old_forward = teacache_forward_fn.__get__(
361
+ transformer, transformer.__class__
362
+ )
363
+ else:
364
+ original_forward = transformer.forward
365
+ transformer.forward = teacache_forward_fn.__get__(
366
+ transformer, transformer.__class__
367
+ )
368
+
369
+ logger.info(
370
+ f"TeaCache enabled for {class_name}: "
371
+ f"steps={num_inference_steps}, threshold={rel_l1_thresh}"
372
+ )
373
+
374
+ try:
375
+ yield pipeline
376
+ finally:
377
+ if has_hook:
378
+ transformer._old_forward = original_forward
379
+ else:
380
+ transformer.forward = original_forward
381
+ logger.debug("TeaCache disabled, original forward restored")
@@ -0,0 +1,99 @@
1
+ {
2
+ "$comment": "TeaCache model registry. Coefficients are 4th-degree polynomial coefficients for np.poly1d rescaling. Source: https://github.com/ali-vilab/TeaCache",
3
+ "models": {
4
+ "flux": {
5
+ "transformer_class": "FluxTransformer2DModel",
6
+ "coefficients": [4.98651651e+02, -2.83781631e+02, 5.58554382e+01, -3.82021401e+00, 2.64230861e-01],
7
+ "default_threshold": 0.6,
8
+ "threshold_guide": "0.25=~1.5x, 0.4=~1.8x, 0.6=~2.0x, 0.8=~2.25x"
9
+ },
10
+ "hunyuan_video": {
11
+ "transformer_class": "HunyuanVideoTransformer3DModel",
12
+ "coefficients": [7.33226126e+02, -4.01131952e+02, 6.75869174e+01, -3.14987800e+00, 9.61237896e-02],
13
+ "default_threshold": 0.15,
14
+ "threshold_guide": "0.1=~1.6x, 0.15=~2.1x"
15
+ },
16
+ "mochi": {
17
+ "transformer_class": "MochiTransformer3DModel",
18
+ "coefficients": [-3.51241319e+03, 8.11675948e+02, -6.09400215e+01, 2.42429681e+00, 3.05291719e-03],
19
+ "default_threshold": 0.09,
20
+ "threshold_guide": "0.06=~1.5x, 0.09=~2.1x"
21
+ },
22
+ "ltx_video": {
23
+ "transformer_class": "LTXVideoTransformer3DModel",
24
+ "coefficients": [2.14700694e+01, -1.28016453e+01, 2.31279151e+00, 7.92487521e-01, 9.69274326e-03],
25
+ "default_threshold": 0.05,
26
+ "threshold_guide": "0.03=~1.6x, 0.05=~2.1x"
27
+ },
28
+ "cogvideox_2b": {
29
+ "transformer_class": "CogVideoXTransformer3DModel",
30
+ "coefficients": [-3.10658903e+01, 2.54732368e+01, -5.92380459e+00, 1.75769064e+00, -3.61568434e-03],
31
+ "default_threshold": 0.1,
32
+ "threshold_guide": "0.1=~1.3x, 0.2=~1.8x"
33
+ },
34
+ "cogvideox_5b": {
35
+ "transformer_class": "CogVideoXTransformer3DModel",
36
+ "coefficients": [-1.53880483e+03, 8.43202495e+02, -1.34363087e+02, 7.97131516e+00, -5.23162339e-02],
37
+ "default_threshold": 0.1,
38
+ "threshold_guide": "0.1=~1.3x, 0.2=~1.8x"
39
+ },
40
+ "cogvideox1.5_5b": {
41
+ "transformer_class": "CogVideoXTransformer3DModel",
42
+ "coefficients": [2.50210439e+02, -1.65061612e+02, 3.57804877e+01, -7.81551492e-01, 3.58559703e-02],
43
+ "default_threshold": 0.2,
44
+ "threshold_guide": "0.1=~1.3x, 0.2=~1.8x, 0.3=~2.1x"
45
+ },
46
+ "cogvideox1.5_5b_i2v": {
47
+ "transformer_class": "CogVideoXTransformer3DModel",
48
+ "coefficients": [1.22842302e+02, -1.04088754e+02, 2.62981677e+01, -3.06009921e-01, 3.71213220e-02],
49
+ "default_threshold": 0.1,
50
+ "threshold_guide": "0.1=~1.5x, 0.2=~2.2x, 0.3=~2.7x"
51
+ },
52
+ "lumina2": {
53
+ "transformer_class": "Lumina2Transformer2DModel",
54
+ "coefficients": [393.76566581, -603.50993606, 209.10239044, -23.00726601, 0.86377344],
55
+ "default_threshold": 0.3,
56
+ "threshold_guide": "0.2=~1.25x, 0.3=~1.56x, 0.4=~2.08x, 0.5=~2.5x"
57
+ },
58
+ "lumina2_v2": {
59
+ "transformer_class": "Lumina2Transformer2DModel",
60
+ "coefficients": [225.7042019806413, -608.8453716535591, 304.1869942338369, 124.21267720116742, -1.4089066892956552],
61
+ "default_threshold": 0.3,
62
+ "threshold_guide": "0.2=~1.5x, 0.3=~1.6x, 0.5=~1.8x, 1.1=~2.1x"
63
+ },
64
+ "wan2.1_t2v_1.3b": {
65
+ "transformer_class": "WanTransformer3DModel",
66
+ "coefficients": [2.39676752e+03, -1.31110545e+03, 2.01331979e+02, -8.29855975e+00, 1.37887774e-01],
67
+ "default_threshold": 0.08,
68
+ "threshold_guide": "0.05=~1.5x, 0.07=~1.6x, 0.08=~2.0x"
69
+ },
70
+ "wan2.1_t2v_14b": {
71
+ "transformer_class": "WanTransformer3DModel",
72
+ "coefficients": [-5784.54975374, 5449.50911966, -1811.16591783, 256.27178429, -13.02252404],
73
+ "default_threshold": 0.2,
74
+ "threshold_guide": "0.14=~1.4x, 0.15=~1.8x, 0.2=~2.0x"
75
+ },
76
+ "wan2.1_i2v_480p": {
77
+ "transformer_class": "WanTransformer3DModel",
78
+ "coefficients": [-3.02331670e+02, 2.23948934e+02, -5.25463970e+01, 5.87348440e+00, -2.01973289e-01],
79
+ "default_threshold": 0.26,
80
+ "threshold_guide": "0.13=~1.6x, 0.19=~2.0x, 0.26=~2.5x"
81
+ },
82
+ "wan2.1_i2v_720p": {
83
+ "transformer_class": "WanTransformer3DModel",
84
+ "coefficients": [-114.36346466, 65.26524496, -18.82220707, 4.91518089, -0.23412683],
85
+ "default_threshold": 0.3,
86
+ "threshold_guide": "0.18=~1.7x, 0.2=~1.9x, 0.3=~2.4x"
87
+ }
88
+ },
89
+ "class_defaults": {
90
+ "$comment": "Default variant used when looking up by transformer class name (no explicit variant specified)",
91
+ "FluxTransformer2DModel": "flux",
92
+ "HunyuanVideoTransformer3DModel": "hunyuan_video",
93
+ "MochiTransformer3DModel": "mochi",
94
+ "LTXVideoTransformer3DModel": "ltx_video",
95
+ "CogVideoXTransformer3DModel": "cogvideox1.5_5b",
96
+ "Lumina2Transformer2DModel": "lumina2",
97
+ "WanTransformer3DModel": "wan2.1_t2v_14b"
98
+ }
99
+ }
dw/test.py ADDED
@@ -0,0 +1,29 @@
1
+ import os
2
+ from .workflow import workflow_from_file
3
+ from . import startup
4
+
5
+
6
+ def main():
7
+ workflow = workflow_from_file(
8
+ os.path.join(
9
+ os.path.dirname(os.path.abspath(__file__)), "workflows", "test.json"
10
+ ),
11
+ "./outputs",
12
+ )
13
+
14
+ try:
15
+ startup("DEBUG")
16
+ workflow.validate()
17
+ except Exception as e:
18
+ print(f"Error validating workflow: {e}")
19
+ exit(1)
20
+
21
+ try:
22
+ workflow.run({})
23
+ except Exception as e:
24
+ print(f"Error running workflow: {e}")
25
+ exit(1)
26
+
27
+
28
+ if __name__ == "__main__":
29
+ main()