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.
- diffusers_workflow-0.4.0a3.dist-info/METADATA +310 -0
- diffusers_workflow-0.4.0a3.dist-info/RECORD +171 -0
- diffusers_workflow-0.4.0a3.dist-info/WHEEL +5 -0
- diffusers_workflow-0.4.0a3.dist-info/entry_points.txt +6 -0
- diffusers_workflow-0.4.0a3.dist-info/licenses/LICENSE +201 -0
- diffusers_workflow-0.4.0a3.dist-info/top_level.txt +1 -0
- dw/__init__.py +353 -0
- dw/arguments.py +906 -0
- dw/cache_blocks.json +16 -0
- dw/cache_blocks.py +145 -0
- dw/community_pipelines/pipeline_flux_rf_inversion.py +1184 -0
- dw/events.py +78 -0
- dw/hub_cache.py +289 -0
- dw/introspection.py +458 -0
- dw/log_setup.py +45 -0
- dw/pipeline_processors/chain.py +750 -0
- dw/pipeline_processors/config_objects.py +235 -0
- dw/pipeline_processors/pipeline.py +1687 -0
- dw/pipeline_processors/remote.py +18 -0
- dw/previous_results.py +259 -0
- dw/prompt_weighting.py +378 -0
- dw/repl.py +298 -0
- dw/repl_commands.py +808 -0
- dw/repl_worker.py +129 -0
- dw/result.py +850 -0
- dw/run.py +92 -0
- dw/schema.py +24 -0
- dw/security.py +379 -0
- dw/serve.py +70 -0
- dw/server/__init__.py +2 -0
- dw/server/app.py +588 -0
- dw/server/jobs.py +547 -0
- dw/server/ui/assets/abap-08VXUWAP.js +1 -0
- dw/server/ui/assets/apex-BWPQTe0t.js +1 -0
- dw/server/ui/assets/azcli-Bc_sGQ0U.js +1 -0
- dw/server/ui/assets/bat-i0X4ZdIN.js +1 -0
- dw/server/ui/assets/bicep-B5-_aFwp.js +2 -0
- dw/server/ui/assets/cameligo-DMUM7wLl.js +1 -0
- dw/server/ui/assets/clojure-Cm7r79vr.js +1 -0
- dw/server/ui/assets/codicon-Brq4_Ui5.ttf +0 -0
- dw/server/ui/assets/coffee-Ba7i2nA0.js +1 -0
- dw/server/ui/assets/cpp-C7h46wYY.js +1 -0
- dw/server/ui/assets/csharp-BKxtCVv1.js +1 -0
- dw/server/ui/assets/csp-bTuwJoIa.js +1 -0
- dw/server/ui/assets/css-DIMkf-bt.js +3 -0
- dw/server/ui/assets/css.worker-B3ciXF_0.js +93 -0
- dw/server/ui/assets/cssMode-CEh6hWi2.js +1 -0
- dw/server/ui/assets/cypher-CVaqCwHa.js +1 -0
- dw/server/ui/assets/dart-onAF5SnQ.js +1 -0
- dw/server/ui/assets/dockerfile-DZFCIeNp.js +1 -0
- dw/server/ui/assets/ecl-D05T4iGw.js +1 -0
- dw/server/ui/assets/editor-jjEx9u7D.css +1 -0
- dw/server/ui/assets/editor.api-CExg3_mM.js +847 -0
- dw/server/ui/assets/editor.worker-q-txB4vs.js +30 -0
- dw/server/ui/assets/elixir-6RTg0lbw.js +1 -0
- dw/server/ui/assets/flow9-C5_-GSwl.js +1 -0
- dw/server/ui/assets/freemarker2-DH6orYh2.js +3 -0
- dw/server/ui/assets/fsharp-C8Ef5oNN.js +1 -0
- dw/server/ui/assets/go-C-y9NEjX.js +1 -0
- dw/server/ui/assets/graphql-fmXr3nnJ.js +1 -0
- dw/server/ui/assets/handlebars-CbrMVW4Q.js +1 -0
- dw/server/ui/assets/hcl-CpzslTdj.js +1 -0
- dw/server/ui/assets/html-YDNPZw2M.js +1 -0
- dw/server/ui/assets/html.worker-C93Ht9o9.js +506 -0
- dw/server/ui/assets/htmlMode-B_zSGWO2.js +1 -0
- dw/server/ui/assets/index-B7-VcYS-.css +1 -0
- dw/server/ui/assets/index-D_EiPU3b.js +13 -0
- dw/server/ui/assets/ini-sBoK_t0W.js +1 -0
- dw/server/ui/assets/java-BEtHBSE6.js +1 -0
- dw/server/ui/assets/javascript-dYuBvioq.js +1 -0
- dw/server/ui/assets/json.worker-B2V3pomh.js +62 -0
- dw/server/ui/assets/jsonMode-CUqLM39V.js +7 -0
- dw/server/ui/assets/julia-Bri6UV-V.js +1 -0
- dw/server/ui/assets/kotlin-BOotOW0E.js +1 -0
- dw/server/ui/assets/less-B9JPFI3C.js +2 -0
- dw/server/ui/assets/lexon-CfSJPG6W.js +1 -0
- dw/server/ui/assets/liquid-D6vxBzMv.js +1 -0
- dw/server/ui/assets/lspLanguageFeatures-1WJ2palX.js +4 -0
- dw/server/ui/assets/lua-CsQS60Ue.js +1 -0
- dw/server/ui/assets/m3-D-oSqn_W.js +1 -0
- dw/server/ui/assets/markdown-Cimd5fb3.js +1 -0
- dw/server/ui/assets/mdx-SHQb6vmD.js +1 -0
- dw/server/ui/assets/mips-CIPQ_RoX.js +1 -0
- dw/server/ui/assets/monaco--ixms01u.css +1 -0
- dw/server/ui/assets/monaco-CP-s5rcP.js +56 -0
- dw/server/ui/assets/msdax-DauUninz.js +1 -0
- dw/server/ui/assets/mysql-SOo6toE5.js +1 -0
- dw/server/ui/assets/objective-c-FvmIjYaQ.js +1 -0
- dw/server/ui/assets/pascal-DrH0SRf2.js +1 -0
- dw/server/ui/assets/pascaligo-D-ptJ9y-.js +1 -0
- dw/server/ui/assets/perl-oz_6vUea.js +1 -0
- dw/server/ui/assets/pgsql-DTj74zXo.js +1 -0
- dw/server/ui/assets/php-nr791fC2.js +1 -0
- dw/server/ui/assets/pla-CopQ2nXW.js +1 -0
- dw/server/ui/assets/postiats-43DmfD33.js +1 -0
- dw/server/ui/assets/powerquery-D3hlyOfw.js +1 -0
- dw/server/ui/assets/powershell-DmHpPYUd.js +1 -0
- dw/server/ui/assets/protobuf-C531GsRP.js +2 -0
- dw/server/ui/assets/pug-Z5eAx3Zn.js +1 -0
- dw/server/ui/assets/python-x0_EGHq9.js +1 -0
- dw/server/ui/assets/qsharp-DkqhCAOL.js +1 -0
- dw/server/ui/assets/r-BwWrilGY.js +1 -0
- dw/server/ui/assets/razor-BZC4LQDP.js +1 -0
- dw/server/ui/assets/redis-ClamHrr6.js +1 -0
- dw/server/ui/assets/redshift-DT7zqm-g.js +1 -0
- dw/server/ui/assets/restructuredtext-BYgofb2h.js +1 -0
- dw/server/ui/assets/ruby-DezsRK8O.js +1 -0
- dw/server/ui/assets/rust-DdL9SqIa.js +1 -0
- dw/server/ui/assets/sb-CcwsVR0C.js +1 -0
- dw/server/ui/assets/scala-DHpiXF5c.js +1 -0
- dw/server/ui/assets/scheme-BeGwcela.js +1 -0
- dw/server/ui/assets/scss-gp-XZpBa.js +3 -0
- dw/server/ui/assets/shell-CC2rA5mh.js +1 -0
- dw/server/ui/assets/solidity-BEEn4gHE.js +1 -0
- dw/server/ui/assets/sophia-CRfGWb83.js +1 -0
- dw/server/ui/assets/sparql-D_Lu-MrJ.js +1 -0
- dw/server/ui/assets/sql-NEE52Syq.js +1 -0
- dw/server/ui/assets/st-DbInun42.js +1 -0
- dw/server/ui/assets/swift-Bxkupp3x.js +1 -0
- dw/server/ui/assets/systemverilog-Bz4Y3fRF.js +1 -0
- dw/server/ui/assets/tcl-DISqw1ZD.js +1 -0
- dw/server/ui/assets/ts.worker-D7T1-Ig5.js +67738 -0
- dw/server/ui/assets/tsMode-BTfA6SbD.js +11 -0
- dw/server/ui/assets/twig-De2hgUGE.js +1 -0
- dw/server/ui/assets/typescript-CWA4MsNk.js +1 -0
- dw/server/ui/assets/typespec-B8J7ngcE.js +1 -0
- dw/server/ui/assets/vb-DV3o63ZY.js +1 -0
- dw/server/ui/assets/wgsl-DpFanUEy.js +298 -0
- dw/server/ui/assets/workers-CWU0uvj5.js +1 -0
- dw/server/ui/assets/xml-KmfTm3rg.js +1 -0
- dw/server/ui/assets/yaml-nFO_dDS6.js +1 -0
- dw/server/ui/index.html +17 -0
- dw/settings.py +77 -0
- dw/step.py +132 -0
- dw/tasks/audio_utils.py +266 -0
- dw/tasks/background_remover.py +43 -0
- dw/tasks/borders.py +113 -0
- dw/tasks/concat_videos.py +80 -0
- dw/tasks/depth_estimator.py +54 -0
- dw/tasks/diffusion_upscale.py +109 -0
- dw/tasks/format_messages.py +24 -0
- dw/tasks/gather.py +139 -0
- dw/tasks/image_to_text.py +43 -0
- dw/tasks/image_utils.py +661 -0
- dw/tasks/interpolate_frames.py +227 -0
- dw/tasks/model_cache.py +39 -0
- dw/tasks/pair_audio.py +58 -0
- dw/tasks/qr_code.py +19 -0
- dw/tasks/restore_faces.py +175 -0
- dw/tasks/rife_model.py +192 -0
- dw/tasks/segment.py +121 -0
- dw/tasks/task.py +474 -0
- dw/tasks/tensor_image.py +57 -0
- dw/tasks/text_generation.py +168 -0
- dw/tasks/text_sections.py +80 -0
- dw/tasks/upscale.py +203 -0
- dw/tasks/video_utils.py +154 -0
- dw/tasks/zoe_depth.py +71 -0
- dw/teacache.py +376 -0
- dw/teacache_models.json +99 -0
- dw/test.py +29 -0
- dw/type_helpers.py +68 -0
- dw/validate.py +43 -0
- dw/variables.py +153 -0
- dw/worker.py +517 -0
- dw/workflow.py +553 -0
- dw/workflow_schema.json +1157 -0
- dw/workflows/augment_prompt.json +65 -0
- dw/workflows/describe_image.json +58 -0
- dw/workflows/h3_context_ir.json +57 -0
- 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")
|
dw/teacache_models.json
ADDED
|
@@ -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()
|