diffusers-workflow 0.4.0a3__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 (171) hide show
  1. diffusers_workflow-0.4.0a3.dist-info/METADATA +310 -0
  2. diffusers_workflow-0.4.0a3.dist-info/RECORD +171 -0
  3. diffusers_workflow-0.4.0a3.dist-info/WHEEL +5 -0
  4. diffusers_workflow-0.4.0a3.dist-info/entry_points.txt +6 -0
  5. diffusers_workflow-0.4.0a3.dist-info/licenses/LICENSE +201 -0
  6. diffusers_workflow-0.4.0a3.dist-info/top_level.txt +1 -0
  7. dw/__init__.py +353 -0
  8. dw/arguments.py +906 -0
  9. dw/cache_blocks.json +16 -0
  10. dw/cache_blocks.py +145 -0
  11. dw/community_pipelines/pipeline_flux_rf_inversion.py +1184 -0
  12. dw/events.py +78 -0
  13. dw/hub_cache.py +289 -0
  14. dw/introspection.py +458 -0
  15. dw/log_setup.py +45 -0
  16. dw/pipeline_processors/chain.py +750 -0
  17. dw/pipeline_processors/config_objects.py +235 -0
  18. dw/pipeline_processors/pipeline.py +1687 -0
  19. dw/pipeline_processors/remote.py +18 -0
  20. dw/previous_results.py +259 -0
  21. dw/prompt_weighting.py +378 -0
  22. dw/repl.py +298 -0
  23. dw/repl_commands.py +808 -0
  24. dw/repl_worker.py +129 -0
  25. dw/result.py +850 -0
  26. dw/run.py +92 -0
  27. dw/schema.py +24 -0
  28. dw/security.py +379 -0
  29. dw/serve.py +70 -0
  30. dw/server/__init__.py +2 -0
  31. dw/server/app.py +588 -0
  32. dw/server/jobs.py +547 -0
  33. dw/server/ui/assets/abap-08VXUWAP.js +1 -0
  34. dw/server/ui/assets/apex-BWPQTe0t.js +1 -0
  35. dw/server/ui/assets/azcli-Bc_sGQ0U.js +1 -0
  36. dw/server/ui/assets/bat-i0X4ZdIN.js +1 -0
  37. dw/server/ui/assets/bicep-B5-_aFwp.js +2 -0
  38. dw/server/ui/assets/cameligo-DMUM7wLl.js +1 -0
  39. dw/server/ui/assets/clojure-Cm7r79vr.js +1 -0
  40. dw/server/ui/assets/codicon-Brq4_Ui5.ttf +0 -0
  41. dw/server/ui/assets/coffee-Ba7i2nA0.js +1 -0
  42. dw/server/ui/assets/cpp-C7h46wYY.js +1 -0
  43. dw/server/ui/assets/csharp-BKxtCVv1.js +1 -0
  44. dw/server/ui/assets/csp-bTuwJoIa.js +1 -0
  45. dw/server/ui/assets/css-DIMkf-bt.js +3 -0
  46. dw/server/ui/assets/css.worker-B3ciXF_0.js +93 -0
  47. dw/server/ui/assets/cssMode-CEh6hWi2.js +1 -0
  48. dw/server/ui/assets/cypher-CVaqCwHa.js +1 -0
  49. dw/server/ui/assets/dart-onAF5SnQ.js +1 -0
  50. dw/server/ui/assets/dockerfile-DZFCIeNp.js +1 -0
  51. dw/server/ui/assets/ecl-D05T4iGw.js +1 -0
  52. dw/server/ui/assets/editor-jjEx9u7D.css +1 -0
  53. dw/server/ui/assets/editor.api-CExg3_mM.js +847 -0
  54. dw/server/ui/assets/editor.worker-q-txB4vs.js +30 -0
  55. dw/server/ui/assets/elixir-6RTg0lbw.js +1 -0
  56. dw/server/ui/assets/flow9-C5_-GSwl.js +1 -0
  57. dw/server/ui/assets/freemarker2-DH6orYh2.js +3 -0
  58. dw/server/ui/assets/fsharp-C8Ef5oNN.js +1 -0
  59. dw/server/ui/assets/go-C-y9NEjX.js +1 -0
  60. dw/server/ui/assets/graphql-fmXr3nnJ.js +1 -0
  61. dw/server/ui/assets/handlebars-CbrMVW4Q.js +1 -0
  62. dw/server/ui/assets/hcl-CpzslTdj.js +1 -0
  63. dw/server/ui/assets/html-YDNPZw2M.js +1 -0
  64. dw/server/ui/assets/html.worker-C93Ht9o9.js +506 -0
  65. dw/server/ui/assets/htmlMode-B_zSGWO2.js +1 -0
  66. dw/server/ui/assets/index-B7-VcYS-.css +1 -0
  67. dw/server/ui/assets/index-D_EiPU3b.js +13 -0
  68. dw/server/ui/assets/ini-sBoK_t0W.js +1 -0
  69. dw/server/ui/assets/java-BEtHBSE6.js +1 -0
  70. dw/server/ui/assets/javascript-dYuBvioq.js +1 -0
  71. dw/server/ui/assets/json.worker-B2V3pomh.js +62 -0
  72. dw/server/ui/assets/jsonMode-CUqLM39V.js +7 -0
  73. dw/server/ui/assets/julia-Bri6UV-V.js +1 -0
  74. dw/server/ui/assets/kotlin-BOotOW0E.js +1 -0
  75. dw/server/ui/assets/less-B9JPFI3C.js +2 -0
  76. dw/server/ui/assets/lexon-CfSJPG6W.js +1 -0
  77. dw/server/ui/assets/liquid-D6vxBzMv.js +1 -0
  78. dw/server/ui/assets/lspLanguageFeatures-1WJ2palX.js +4 -0
  79. dw/server/ui/assets/lua-CsQS60Ue.js +1 -0
  80. dw/server/ui/assets/m3-D-oSqn_W.js +1 -0
  81. dw/server/ui/assets/markdown-Cimd5fb3.js +1 -0
  82. dw/server/ui/assets/mdx-SHQb6vmD.js +1 -0
  83. dw/server/ui/assets/mips-CIPQ_RoX.js +1 -0
  84. dw/server/ui/assets/monaco--ixms01u.css +1 -0
  85. dw/server/ui/assets/monaco-CP-s5rcP.js +56 -0
  86. dw/server/ui/assets/msdax-DauUninz.js +1 -0
  87. dw/server/ui/assets/mysql-SOo6toE5.js +1 -0
  88. dw/server/ui/assets/objective-c-FvmIjYaQ.js +1 -0
  89. dw/server/ui/assets/pascal-DrH0SRf2.js +1 -0
  90. dw/server/ui/assets/pascaligo-D-ptJ9y-.js +1 -0
  91. dw/server/ui/assets/perl-oz_6vUea.js +1 -0
  92. dw/server/ui/assets/pgsql-DTj74zXo.js +1 -0
  93. dw/server/ui/assets/php-nr791fC2.js +1 -0
  94. dw/server/ui/assets/pla-CopQ2nXW.js +1 -0
  95. dw/server/ui/assets/postiats-43DmfD33.js +1 -0
  96. dw/server/ui/assets/powerquery-D3hlyOfw.js +1 -0
  97. dw/server/ui/assets/powershell-DmHpPYUd.js +1 -0
  98. dw/server/ui/assets/protobuf-C531GsRP.js +2 -0
  99. dw/server/ui/assets/pug-Z5eAx3Zn.js +1 -0
  100. dw/server/ui/assets/python-x0_EGHq9.js +1 -0
  101. dw/server/ui/assets/qsharp-DkqhCAOL.js +1 -0
  102. dw/server/ui/assets/r-BwWrilGY.js +1 -0
  103. dw/server/ui/assets/razor-BZC4LQDP.js +1 -0
  104. dw/server/ui/assets/redis-ClamHrr6.js +1 -0
  105. dw/server/ui/assets/redshift-DT7zqm-g.js +1 -0
  106. dw/server/ui/assets/restructuredtext-BYgofb2h.js +1 -0
  107. dw/server/ui/assets/ruby-DezsRK8O.js +1 -0
  108. dw/server/ui/assets/rust-DdL9SqIa.js +1 -0
  109. dw/server/ui/assets/sb-CcwsVR0C.js +1 -0
  110. dw/server/ui/assets/scala-DHpiXF5c.js +1 -0
  111. dw/server/ui/assets/scheme-BeGwcela.js +1 -0
  112. dw/server/ui/assets/scss-gp-XZpBa.js +3 -0
  113. dw/server/ui/assets/shell-CC2rA5mh.js +1 -0
  114. dw/server/ui/assets/solidity-BEEn4gHE.js +1 -0
  115. dw/server/ui/assets/sophia-CRfGWb83.js +1 -0
  116. dw/server/ui/assets/sparql-D_Lu-MrJ.js +1 -0
  117. dw/server/ui/assets/sql-NEE52Syq.js +1 -0
  118. dw/server/ui/assets/st-DbInun42.js +1 -0
  119. dw/server/ui/assets/swift-Bxkupp3x.js +1 -0
  120. dw/server/ui/assets/systemverilog-Bz4Y3fRF.js +1 -0
  121. dw/server/ui/assets/tcl-DISqw1ZD.js +1 -0
  122. dw/server/ui/assets/ts.worker-D7T1-Ig5.js +67738 -0
  123. dw/server/ui/assets/tsMode-BTfA6SbD.js +11 -0
  124. dw/server/ui/assets/twig-De2hgUGE.js +1 -0
  125. dw/server/ui/assets/typescript-CWA4MsNk.js +1 -0
  126. dw/server/ui/assets/typespec-B8J7ngcE.js +1 -0
  127. dw/server/ui/assets/vb-DV3o63ZY.js +1 -0
  128. dw/server/ui/assets/wgsl-DpFanUEy.js +298 -0
  129. dw/server/ui/assets/workers-CWU0uvj5.js +1 -0
  130. dw/server/ui/assets/xml-KmfTm3rg.js +1 -0
  131. dw/server/ui/assets/yaml-nFO_dDS6.js +1 -0
  132. dw/server/ui/index.html +17 -0
  133. dw/settings.py +77 -0
  134. dw/step.py +132 -0
  135. dw/tasks/audio_utils.py +266 -0
  136. dw/tasks/background_remover.py +43 -0
  137. dw/tasks/borders.py +113 -0
  138. dw/tasks/concat_videos.py +80 -0
  139. dw/tasks/depth_estimator.py +54 -0
  140. dw/tasks/diffusion_upscale.py +109 -0
  141. dw/tasks/format_messages.py +24 -0
  142. dw/tasks/gather.py +139 -0
  143. dw/tasks/image_to_text.py +43 -0
  144. dw/tasks/image_utils.py +661 -0
  145. dw/tasks/interpolate_frames.py +227 -0
  146. dw/tasks/model_cache.py +39 -0
  147. dw/tasks/pair_audio.py +58 -0
  148. dw/tasks/qr_code.py +19 -0
  149. dw/tasks/restore_faces.py +175 -0
  150. dw/tasks/rife_model.py +192 -0
  151. dw/tasks/segment.py +121 -0
  152. dw/tasks/task.py +474 -0
  153. dw/tasks/tensor_image.py +57 -0
  154. dw/tasks/text_generation.py +168 -0
  155. dw/tasks/text_sections.py +80 -0
  156. dw/tasks/upscale.py +203 -0
  157. dw/tasks/video_utils.py +154 -0
  158. dw/tasks/zoe_depth.py +71 -0
  159. dw/teacache.py +376 -0
  160. dw/teacache_models.json +99 -0
  161. dw/test.py +29 -0
  162. dw/type_helpers.py +68 -0
  163. dw/validate.py +43 -0
  164. dw/variables.py +153 -0
  165. dw/worker.py +517 -0
  166. dw/workflow.py +553 -0
  167. dw/workflow_schema.json +1157 -0
  168. dw/workflows/augment_prompt.json +65 -0
  169. dw/workflows/describe_image.json +58 -0
  170. dw/workflows/h3_context_ir.json +57 -0
  171. dw/workflows/test.json +31 -0
dw/teacache.py ADDED
@@ -0,0 +1,376 @@
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 cnt, accumulated_rel_l1_distance, previous_modulated_input, previous_residual, previous_timestep
114
+
115
+ # TeaCache assumes exactly one transformer forward call per denoising
116
+ # step. Pipelines running true classifier-free guidance (e.g. Flux with
117
+ # negative_prompt + true_cfg_scale > 1) call the transformer twice per
118
+ # step -- once for the conditional pass and once for the unconditional
119
+ # pass -- using the identical timestep both times. That second call
120
+ # would silently share/corrupt previous_modulated_input and
121
+ # previous_residual across the two passes, so detect it and fail loudly
122
+ # instead of producing a corrupted image.
123
+ if (
124
+ timestep is not None
125
+ and previous_timestep is not None
126
+ and timestep.shape == previous_timestep.shape
127
+ and torch.equal(timestep, previous_timestep)
128
+ ):
129
+ raise RuntimeError(
130
+ "TeaCache does not support true classifier-free guidance "
131
+ "(negative_prompt with true_cfg_scale > 1); disable one of them. "
132
+ "Detected two transformer forward calls with an identical "
133
+ "timestep within a single denoising step, which would corrupt "
134
+ "TeaCache's cached state."
135
+ )
136
+ if timestep is not None:
137
+ previous_timestep = timestep.detach().clone()
138
+
139
+ if joint_attention_kwargs is not None:
140
+ joint_attention_kwargs = joint_attention_kwargs.copy()
141
+ lora_scale = joint_attention_kwargs.pop("scale", 1.0)
142
+ else:
143
+ lora_scale = 1.0
144
+
145
+ if USE_PEFT_BACKEND:
146
+ scale_lora_layers(self, lora_scale)
147
+
148
+ hidden_states = self.x_embedder(hidden_states)
149
+
150
+ timestep = timestep.to(hidden_states.dtype) * 1000
151
+ if guidance is not None:
152
+ guidance = guidance.to(hidden_states.dtype) * 1000
153
+
154
+ temb = (
155
+ self.time_text_embed(timestep, pooled_projections)
156
+ if guidance is None
157
+ else self.time_text_embed(timestep, guidance, pooled_projections)
158
+ )
159
+ encoder_hidden_states = self.context_embedder(encoder_hidden_states)
160
+
161
+ if txt_ids.ndim == 3:
162
+ txt_ids = txt_ids[0]
163
+ if img_ids.ndim == 3:
164
+ img_ids = img_ids[0]
165
+
166
+ ids = torch.cat((txt_ids, img_ids), dim=0)
167
+ image_rotary_emb = self.pos_embed(ids)
168
+
169
+ if (
170
+ joint_attention_kwargs is not None
171
+ and "ip_adapter_image_embeds" in joint_attention_kwargs
172
+ ):
173
+ ip_adapter_image_embeds = joint_attention_kwargs.pop(
174
+ "ip_adapter_image_embeds"
175
+ )
176
+ ip_hidden_states = self.encoder_hid_proj(ip_adapter_image_embeds)
177
+ joint_attention_kwargs.update({"ip_hidden_states": ip_hidden_states})
178
+
179
+ # TeaCache: extract cache signal from first block's normalization
180
+ inp = hidden_states.clone()
181
+ temb_ = temb.clone()
182
+ modulated_inp, gate_msa, shift_mlp, scale_mlp, gate_mlp = (
183
+ self.transformer_blocks[0].norm1(inp, emb=temb_)
184
+ )
185
+
186
+ # Decide whether to compute or reuse cached result
187
+ if cnt == 0 or cnt == num_inference_steps - 1:
188
+ should_calc = True
189
+ accumulated_rel_l1_distance = 0
190
+ else:
191
+ relative_diff = (
192
+ (
193
+ (modulated_inp - previous_modulated_input).abs().mean()
194
+ / previous_modulated_input.abs().mean()
195
+ )
196
+ .cpu()
197
+ .item()
198
+ )
199
+ accumulated_rel_l1_distance += rescale_func(relative_diff)
200
+
201
+ if accumulated_rel_l1_distance < rel_l1_thresh:
202
+ should_calc = False
203
+ else:
204
+ should_calc = True
205
+ accumulated_rel_l1_distance = 0
206
+
207
+ previous_modulated_input = modulated_inp
208
+ cnt += 1
209
+ if cnt == num_inference_steps:
210
+ cnt = 0
211
+
212
+ if not should_calc:
213
+ hidden_states += previous_residual
214
+ else:
215
+ ori_hidden_states = hidden_states.clone()
216
+
217
+ # No gradient-checkpointing branch: this forward only runs under
218
+ # Pipeline.run's @torch.inference_mode(), so grads are never enabled
219
+ for index_block, block in enumerate(self.transformer_blocks):
220
+ encoder_hidden_states, hidden_states = block(
221
+ hidden_states=hidden_states,
222
+ encoder_hidden_states=encoder_hidden_states,
223
+ temb=temb,
224
+ image_rotary_emb=image_rotary_emb,
225
+ joint_attention_kwargs=joint_attention_kwargs,
226
+ )
227
+
228
+ if controlnet_block_samples is not None:
229
+ interval_control = len(self.transformer_blocks) / len(
230
+ controlnet_block_samples
231
+ )
232
+ interval_control = int(np.ceil(interval_control))
233
+ if controlnet_blocks_repeat:
234
+ hidden_states = (
235
+ hidden_states
236
+ + controlnet_block_samples[
237
+ index_block % len(controlnet_block_samples)
238
+ ]
239
+ )
240
+ else:
241
+ hidden_states = (
242
+ hidden_states
243
+ + controlnet_block_samples[index_block // interval_control]
244
+ )
245
+
246
+ for index_block, block in enumerate(self.single_transformer_blocks):
247
+ encoder_hidden_states, hidden_states = block(
248
+ hidden_states=hidden_states,
249
+ encoder_hidden_states=encoder_hidden_states,
250
+ temb=temb,
251
+ image_rotary_emb=image_rotary_emb,
252
+ joint_attention_kwargs=joint_attention_kwargs,
253
+ )
254
+
255
+ if controlnet_single_block_samples is not None:
256
+ interval_control = len(self.single_transformer_blocks) / len(
257
+ controlnet_single_block_samples
258
+ )
259
+ interval_control = int(np.ceil(interval_control))
260
+ hidden_states[:, encoder_hidden_states.shape[1] :, ...] = (
261
+ hidden_states[:, encoder_hidden_states.shape[1] :, ...]
262
+ + controlnet_single_block_samples[
263
+ index_block // interval_control
264
+ ]
265
+ )
266
+
267
+ previous_residual = hidden_states - ori_hidden_states
268
+
269
+ hidden_states = self.norm_out(hidden_states, temb)
270
+ output = self.proj_out(hidden_states)
271
+
272
+ if USE_PEFT_BACKEND:
273
+ unscale_lora_layers(self, lora_scale)
274
+
275
+ if not return_dict:
276
+ return (output,)
277
+
278
+ return Transformer2DModelOutput(sample=output)
279
+
280
+ return teacache_forward
281
+
282
+
283
+ # Map transformer class names to their forward factory functions.
284
+ # Models in the JSON registry without a factory here will get an informative error.
285
+ _FORWARD_FACTORIES = {
286
+ "FluxTransformer2DModel": _create_flux_teacache_forward,
287
+ }
288
+
289
+
290
+ # ---------------------------------------------------------------------------
291
+ # Public API
292
+ # ---------------------------------------------------------------------------
293
+
294
+
295
+ @contextmanager
296
+ def teacache_context(
297
+ pipeline, num_inference_steps, rel_l1_thresh=None, coefficients=None, variant=None
298
+ ):
299
+ """Context manager that enables TeaCache on a pipeline's transformer.
300
+
301
+ Auto-detects the transformer type and applies the appropriate
302
+ TeaCache forward function. Restores original forward on exit.
303
+
304
+ Args:
305
+ pipeline: A DiffusionPipeline with a .transformer attribute
306
+ num_inference_steps: Number of inference steps (must match pipeline call)
307
+ rel_l1_thresh: Cache threshold override. If None, uses model default.
308
+ Higher = more speedup, more quality loss.
309
+ coefficients: Polynomial coefficients override. If None, uses model default.
310
+ List of 5 floats for np.poly1d rescaling function.
311
+ variant: Explicit model variant name (e.g., "wan2.1_t2v_1.3b").
312
+ Required when a transformer class has multiple variants (CogVideoX, Wan).
313
+ If None, uses class_defaults from the registry.
314
+ """
315
+ transformer = pipeline.transformer
316
+ class_name = transformer.__class__.__name__
317
+
318
+ # Look up model info from registry
319
+ model_info = _get_model_info(transformer, variant)
320
+
321
+ # Check we have a forward implementation for this class
322
+ factory = _FORWARD_FACTORIES.get(class_name)
323
+ if factory is None:
324
+ supported = ", ".join(_FORWARD_FACTORIES.keys())
325
+ raise ValueError(
326
+ f"No TeaCache forward implementation for {class_name}. "
327
+ f"Implemented: {supported}. "
328
+ f"The model is in the registry but needs a custom forward function."
329
+ )
330
+
331
+ # Use overrides or defaults
332
+ if rel_l1_thresh is None:
333
+ rel_l1_thresh = model_info["default_threshold"]
334
+ if coefficients is None:
335
+ coefficients = model_info["coefficients"]
336
+
337
+ teacache_forward_fn = factory(num_inference_steps, rel_l1_thresh, coefficients)
338
+
339
+ # accelerate's enable_model_cpu_offload/enable_sequential_cpu_offload installs
340
+ # an AlignDevicesHook via add_hook_to_module (accelerate/hooks.py), which
341
+ # replaces transformer.forward with a wrapper closing over module and the
342
+ # true original forward (stashed as transformer._old_forward, still bound to
343
+ # the instance). That wrapper is what moves the module's weights to the
344
+ # execution device in its pre_forward before calling _old_forward. If we
345
+ # clobber transformer.forward like the no-hook path below, we remove that
346
+ # wrapper entirely: the CPU-resident module then receives CUDA inputs and
347
+ # raises a device-mismatch RuntimeError. Instead, when a hook is present we
348
+ # wrap what the hook considers "the real forward" -- _old_forward -- so the
349
+ # call chain stays hook.forward -> pre_forward (places weights) ->
350
+ # teacache_forward -> post_forward.
351
+ has_hook = hasattr(transformer, "_hf_hook") and hasattr(transformer, "_old_forward")
352
+
353
+ if has_hook:
354
+ original_forward = transformer._old_forward
355
+ transformer._old_forward = teacache_forward_fn.__get__(
356
+ transformer, transformer.__class__
357
+ )
358
+ else:
359
+ original_forward = transformer.forward
360
+ transformer.forward = teacache_forward_fn.__get__(
361
+ transformer, transformer.__class__
362
+ )
363
+
364
+ logger.info(
365
+ f"TeaCache enabled for {class_name}: "
366
+ f"steps={num_inference_steps}, threshold={rel_l1_thresh}"
367
+ )
368
+
369
+ try:
370
+ yield pipeline
371
+ finally:
372
+ if has_hook:
373
+ transformer._old_forward = original_forward
374
+ else:
375
+ transformer.forward = original_forward
376
+ 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()
dw/type_helpers.py ADDED
@@ -0,0 +1,68 @@
1
+ import importlib
2
+
3
+
4
+ def get_type(module_name, type_name):
5
+ module = __import__(module_name)
6
+ return getattr(module, type_name)
7
+
8
+
9
+ def load_type_from_name(type_name):
10
+ if "." in type_name:
11
+ return load_type_from_full_name(type_name)
12
+
13
+ return get_type("diffusers", type_name)
14
+
15
+
16
+ def load_type_from_full_name(full_name):
17
+ # Split the full name into module path and object name
18
+ module_path, object_name = full_name.rsplit(".", 1)
19
+
20
+ # Dynamically import the module
21
+ module = importlib.import_module(module_path)
22
+
23
+ # Get the object from the module
24
+ return getattr(module, object_name)
25
+
26
+
27
+ def has_method(o, name):
28
+ return callable(getattr(o, name, None))
29
+
30
+
31
+ def load_constant_from_name(name):
32
+ """Load a constant declared in python, by its dotted name.
33
+
34
+ The leading run of names that imports is the module the constant lives in and
35
+ the rest are read from it, so a constant held in a dataclass is reachable
36
+ ('...utils.GEMMA4_PROMPT_ENHANCEMENT_CONFIG.max_new_tokens') as well as one
37
+ declared at module scope. A bare name is read from diffusers, matching the way
38
+ a bare type reference resolves.
39
+
40
+ Args:
41
+ name: Dotted name of the constant
42
+
43
+ Returns:
44
+ The value the name refers to
45
+
46
+ Raises:
47
+ ImportError: If no leading part of the name names a module
48
+ AttributeError: If the module has no such attribute
49
+ """
50
+ parts = name.split(".")
51
+
52
+ module, attributes = None, parts
53
+ for i in range(len(parts) - 1, 0, -1):
54
+ try:
55
+ module = importlib.import_module(".".join(parts[:i]))
56
+ attributes = parts[i:]
57
+ break
58
+ except ImportError:
59
+ continue
60
+
61
+ if module is None:
62
+ # No dotted module path - a bare name, read from diffusers
63
+ module = importlib.import_module("diffusers")
64
+
65
+ value = module
66
+ for attribute in attributes:
67
+ value = getattr(value, attribute)
68
+ return value
dw/validate.py ADDED
@@ -0,0 +1,43 @@
1
+ import argparse
2
+ import os
3
+ from .workflow import workflow_from_file
4
+ from . import startup
5
+ from .security import validate_workflow_path, SecurityError
6
+
7
+
8
+ def main():
9
+ parser = argparse.ArgumentParser(description="Validate a workflow from a file.")
10
+ parser.add_argument(
11
+ "file_name", type=str, help="The filespec of the workflow to validate"
12
+ )
13
+
14
+ parser.add_argument(
15
+ "-l",
16
+ "--log_level",
17
+ type=str,
18
+ default="INFO",
19
+ help="Set the logging level (DEBUG, INFO, WARNING, ERROR, CRITICAL)",
20
+ )
21
+ args = parser.parse_args()
22
+
23
+ try:
24
+ validated_file_path = validate_workflow_path(args.file_name)
25
+ if not os.path.exists(validated_file_path):
26
+ raise FileNotFoundError(f"File {validated_file_path} does not exist")
27
+ except SecurityError as e:
28
+ print(f"Error: Security validation failed: {e}")
29
+ exit(1)
30
+
31
+ startup(args.log_level)
32
+
33
+ try:
34
+ workflow = workflow_from_file(validated_file_path, ".")
35
+ workflow.validate()
36
+ print("Workflow validated successfully")
37
+ except Exception as e:
38
+ print(f"Error validating workflow '{args.file_name}': {e}")
39
+ exit(1)
40
+
41
+
42
+ if __name__ == "__main__":
43
+ main()