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
|
@@ -0,0 +1,1687 @@
|
|
|
1
|
+
import torch
|
|
2
|
+
import contextlib
|
|
3
|
+
import copy
|
|
4
|
+
import functools
|
|
5
|
+
import gc
|
|
6
|
+
import importlib
|
|
7
|
+
import inspect
|
|
8
|
+
import logging
|
|
9
|
+
from .config_objects import (
|
|
10
|
+
get_quantization_configuration,
|
|
11
|
+
get_group_offload_configuration,
|
|
12
|
+
get_cache_configuration,
|
|
13
|
+
get_load_components_arguments,
|
|
14
|
+
)
|
|
15
|
+
from .remote import remote_text_encoder
|
|
16
|
+
from ..cache_blocks import register_cache_blocks
|
|
17
|
+
from ..teacache import teacache_context
|
|
18
|
+
from ..type_helpers import has_method
|
|
19
|
+
from .. import empty_device_cache, get_device_type
|
|
20
|
+
from diffusers import attention_backend
|
|
21
|
+
|
|
22
|
+
# dw.prompt_weighting (transformers) and diffusers.hooks (peft, bitsandbytes) are
|
|
23
|
+
# imported where they are used - at module scope they add seconds to every startup
|
|
24
|
+
|
|
25
|
+
from ..events import WorkflowCancelled, get_context
|
|
26
|
+
|
|
27
|
+
logger = logging.getLogger("dw")
|
|
28
|
+
|
|
29
|
+
# Names the cache state a run accumulates. Any stable string works - it only has to
|
|
30
|
+
# match itself across the steps of one call
|
|
31
|
+
_CACHE_CONTEXT_NAME = "dw"
|
|
32
|
+
|
|
33
|
+
optional_component_names = [
|
|
34
|
+
"controlnet",
|
|
35
|
+
"transformer",
|
|
36
|
+
"transformer_2",
|
|
37
|
+
"vae",
|
|
38
|
+
"unet",
|
|
39
|
+
"text_encoder",
|
|
40
|
+
"text_encoder_2",
|
|
41
|
+
"text_encoder_3",
|
|
42
|
+
"tokenizer",
|
|
43
|
+
"tokenizer_2",
|
|
44
|
+
"tokenizer_3",
|
|
45
|
+
"image_encoder",
|
|
46
|
+
"feature_extractor",
|
|
47
|
+
"prompt_enhancer_head",
|
|
48
|
+
"model",
|
|
49
|
+
]
|
|
50
|
+
|
|
51
|
+
# Pipeline-definition keys that can never name a component
|
|
52
|
+
_NON_COMPONENT_KEYS = {
|
|
53
|
+
"configuration",
|
|
54
|
+
"from_pretrained_arguments",
|
|
55
|
+
"arguments",
|
|
56
|
+
"scheduler",
|
|
57
|
+
"loras",
|
|
58
|
+
"ip_adapter",
|
|
59
|
+
"seed",
|
|
60
|
+
"remote_text_encoder",
|
|
61
|
+
}
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
def declared_component_names(pipeline_definition):
|
|
65
|
+
"""The component names a pipeline definition can load or configure.
|
|
66
|
+
|
|
67
|
+
The known names plus any other key shaped like a component - a dict carrying
|
|
68
|
+
'from_pretrained_arguments' (a scheduler carries 'from_config_args' instead).
|
|
69
|
+
Diffusers grows new component names faster than the list above; a workflow
|
|
70
|
+
naming one gets it loaded rather than silently dropped.
|
|
71
|
+
"""
|
|
72
|
+
names = list(optional_component_names)
|
|
73
|
+
for key, value in pipeline_definition.items():
|
|
74
|
+
if (
|
|
75
|
+
key not in names
|
|
76
|
+
and key not in _NON_COMPONENT_KEYS
|
|
77
|
+
and isinstance(value, dict)
|
|
78
|
+
and "from_pretrained_arguments" in value
|
|
79
|
+
):
|
|
80
|
+
logger.info(f"Treating '{key}' as a component definition")
|
|
81
|
+
names.append(key)
|
|
82
|
+
return names
|
|
83
|
+
|
|
84
|
+
|
|
85
|
+
class Pipeline:
|
|
86
|
+
"""
|
|
87
|
+
Manages pipeline initialization, configuration, and execution.
|
|
88
|
+
Handles loading of models, schedulers, and adapters.
|
|
89
|
+
"""
|
|
90
|
+
|
|
91
|
+
def __init__(
|
|
92
|
+
self,
|
|
93
|
+
pipeline_definition,
|
|
94
|
+
default_seed,
|
|
95
|
+
device,
|
|
96
|
+
pipeline=None,
|
|
97
|
+
output_dir=None,
|
|
98
|
+
file_prefix=None,
|
|
99
|
+
):
|
|
100
|
+
"""
|
|
101
|
+
Initialize pipeline with configuration and device settings.
|
|
102
|
+
|
|
103
|
+
Args:
|
|
104
|
+
pipeline_definition: Dictionary containing pipeline configuration
|
|
105
|
+
default_seed: Seed value for reproducibility
|
|
106
|
+
device: Device to run pipeline on (e.g., 'cuda', 'mps', 'cpu') - the
|
|
107
|
+
configuration's own 'device' takes precedence over it
|
|
108
|
+
pipeline: Optional existing pipeline to use
|
|
109
|
+
output_dir: The workflow's output directory - where a chained run
|
|
110
|
+
with save_segments writes its segment files
|
|
111
|
+
file_prefix: Naming prefix for those files, matching the step's
|
|
112
|
+
result naming (workflow id + step name)
|
|
113
|
+
"""
|
|
114
|
+
self.pipeline_definition = pipeline_definition
|
|
115
|
+
self.default_seed = default_seed
|
|
116
|
+
# A step can pin itself to a device, overriding the one dw is running on. It
|
|
117
|
+
# becomes the default for this pipeline's components as well
|
|
118
|
+
self.device = self.configuration.get("device", device)
|
|
119
|
+
self.pipeline = pipeline
|
|
120
|
+
self.output_dir = output_dir
|
|
121
|
+
self.file_prefix = file_prefix
|
|
122
|
+
logger.debug(f"Initialized pipeline with device: {self.device}")
|
|
123
|
+
|
|
124
|
+
@property
|
|
125
|
+
def configuration(self):
|
|
126
|
+
return self.pipeline_definition.get("configuration", {})
|
|
127
|
+
|
|
128
|
+
@property
|
|
129
|
+
def name(self):
|
|
130
|
+
return self.from_pretrained_arguments.get("model_name", "")
|
|
131
|
+
|
|
132
|
+
@property
|
|
133
|
+
def from_pretrained_arguments(self):
|
|
134
|
+
return self.pipeline_definition.get("from_pretrained_arguments", {})
|
|
135
|
+
|
|
136
|
+
@property
|
|
137
|
+
def argument_template(self):
|
|
138
|
+
return self.pipeline_definition["arguments"]
|
|
139
|
+
|
|
140
|
+
def component_names(self, key):
|
|
141
|
+
"""The component names one of the sharing lists holds.
|
|
142
|
+
|
|
143
|
+
The lists were only ever read off the pipeline itself, while the schema and
|
|
144
|
+
the guide put them in its configuration - a workflow written to the docs
|
|
145
|
+
shared nothing and said nothing about it. Both places are read now.
|
|
146
|
+
|
|
147
|
+
Args:
|
|
148
|
+
key: 'shared_components' or 'reused_components'
|
|
149
|
+
|
|
150
|
+
Returns:
|
|
151
|
+
List of component names
|
|
152
|
+
"""
|
|
153
|
+
return list(self.pipeline_definition.get(key, [])) + list(
|
|
154
|
+
self.configuration.get(key, [])
|
|
155
|
+
)
|
|
156
|
+
|
|
157
|
+
def resolve_reused_components(self, shared_components):
|
|
158
|
+
"""The components an earlier step shared that this one asks to reuse.
|
|
159
|
+
|
|
160
|
+
Args:
|
|
161
|
+
shared_components: Dictionary of components shared between pipelines
|
|
162
|
+
|
|
163
|
+
Returns:
|
|
164
|
+
Dict of component name to the component itself
|
|
165
|
+
|
|
166
|
+
Raises:
|
|
167
|
+
ValueError: If a name was never shared by an earlier step
|
|
168
|
+
"""
|
|
169
|
+
reused = {}
|
|
170
|
+
for name in self.component_names("reused_components"):
|
|
171
|
+
if name not in shared_components:
|
|
172
|
+
raise ValueError(
|
|
173
|
+
f"Cannot reuse component '{name}' - no earlier step shared it. "
|
|
174
|
+
f"Shared so far: {sorted(shared_components) or 'nothing'}"
|
|
175
|
+
)
|
|
176
|
+
logger.debug(f"Reusing component: {name}")
|
|
177
|
+
reused[name] = shared_components[name]
|
|
178
|
+
return reused
|
|
179
|
+
|
|
180
|
+
def populate_from_pretrained_arguments(self, device, shared_components):
|
|
181
|
+
"""
|
|
182
|
+
Prepare arguments for pipeline creation, including shared components.
|
|
183
|
+
|
|
184
|
+
The loaded components go into a copy, not into the definition they were read
|
|
185
|
+
from. The definition belongs to the workflow and outlives every step, so a
|
|
186
|
+
component stored there is a component the run holds until it ends: releasing
|
|
187
|
+
the pipeline frees nothing, and a workflow that loads a second large model
|
|
188
|
+
after releasing the first runs out of memory holding both. Copying also
|
|
189
|
+
leaves the definition intact for a second load - load_component consumes
|
|
190
|
+
'model_name' out of the arguments it is handed.
|
|
191
|
+
|
|
192
|
+
Args:
|
|
193
|
+
device: Device to run pipeline on
|
|
194
|
+
shared_components: Dictionary of components shared between pipelines
|
|
195
|
+
"""
|
|
196
|
+
logger.debug("Populating from_pretrained arguments")
|
|
197
|
+
from_pretrained_arguments = dict(self.from_pretrained_arguments)
|
|
198
|
+
|
|
199
|
+
# Load optional components (controlnet, vae, unet, etc.), including any
|
|
200
|
+
# component-shaped key outside the known names
|
|
201
|
+
for component_name in declared_component_names(self.pipeline_definition):
|
|
202
|
+
self.load_optional_component(
|
|
203
|
+
component_name, from_pretrained_arguments, device
|
|
204
|
+
)
|
|
205
|
+
|
|
206
|
+
# Handle remote text encoder configuration by setting local text_encoder to None
|
|
207
|
+
if self.pipeline_definition.get("remote_text_encoder", None):
|
|
208
|
+
logger.info("Configuring remote text encoder")
|
|
209
|
+
from_pretrained_arguments["text_encoder"] = None
|
|
210
|
+
|
|
211
|
+
return from_pretrained_arguments
|
|
212
|
+
|
|
213
|
+
def load(self, shared_components):
|
|
214
|
+
"""
|
|
215
|
+
Load and configure the pipeline with all components.
|
|
216
|
+
|
|
217
|
+
Args:
|
|
218
|
+
shared_components: Dictionary of components shared between pipelines
|
|
219
|
+
"""
|
|
220
|
+
logger.debug(f"Loading pipeline: {self.name}")
|
|
221
|
+
|
|
222
|
+
# Import modules that need to register with diffusers/transformers before loading
|
|
223
|
+
# (e.g., sdnq registers its quantization method on import)
|
|
224
|
+
for module_name in self.configuration.get("pre_load_modules", []):
|
|
225
|
+
logger.info(f"Pre-loading module: {module_name}")
|
|
226
|
+
importlib.import_module(module_name)
|
|
227
|
+
|
|
228
|
+
# Prepare arguments and load pipeline
|
|
229
|
+
from_pretrained_arguments = self.populate_from_pretrained_arguments(
|
|
230
|
+
self.device, shared_components
|
|
231
|
+
)
|
|
232
|
+
reused_components = self.resolve_reused_components(shared_components)
|
|
233
|
+
|
|
234
|
+
# Adapters add weights to the components they attach to, and an offloading
|
|
235
|
+
# hook only streams the weights that existed when it was installed - so a
|
|
236
|
+
# pipeline that loads any is placed after they are on it, not at load
|
|
237
|
+
adapters_to_load = bool(self.pipeline_definition.get("loras", [])) or (
|
|
238
|
+
self.pipeline_definition.get("ip_adapter", None) is not None
|
|
239
|
+
)
|
|
240
|
+
|
|
241
|
+
# Load and configure the main pipeline
|
|
242
|
+
self.pipeline = load_component(
|
|
243
|
+
"pipeline",
|
|
244
|
+
self.configuration,
|
|
245
|
+
from_pretrained_arguments,
|
|
246
|
+
self.device,
|
|
247
|
+
reused_components,
|
|
248
|
+
defer_placement=adapters_to_load,
|
|
249
|
+
)
|
|
250
|
+
|
|
251
|
+
# Enable attention slicing if explicitly requested or automatically on MPS
|
|
252
|
+
# MPS benefits from slicing since Metal shares system RAM with the GPU
|
|
253
|
+
if self.configuration.get("enable_attention_slicing", False) or (
|
|
254
|
+
get_device_type(self.device) == "mps"
|
|
255
|
+
and not self.configuration.get("disable_attention_slicing", False)
|
|
256
|
+
):
|
|
257
|
+
# Modular pipelines have no attention slicing - on MPS this is applied
|
|
258
|
+
# automatically, so skip rather than fail when the pipeline lacks it
|
|
259
|
+
if has_method(self.pipeline, "enable_attention_slicing"):
|
|
260
|
+
logger.debug("Enabling attention slicing for pipeline")
|
|
261
|
+
self.pipeline.enable_attention_slicing()
|
|
262
|
+
else:
|
|
263
|
+
logger.debug(
|
|
264
|
+
f"{type(self.pipeline).__name__} does not support attention slicing, skipping"
|
|
265
|
+
)
|
|
266
|
+
|
|
267
|
+
# configure components that are not shared
|
|
268
|
+
self.configure_loaded_components()
|
|
269
|
+
|
|
270
|
+
# Apply SDNQ quantized matmul optimization to specified components
|
|
271
|
+
sdnq_optimize = self.configuration.get("sdnq_optimize", [])
|
|
272
|
+
if sdnq_optimize:
|
|
273
|
+
apply_sdnq_optimizations(self.pipeline, sdnq_optimize)
|
|
274
|
+
|
|
275
|
+
# Enable diffusers built-in cache acceleration on transformer
|
|
276
|
+
cache_config = get_cache_configuration(self.configuration)
|
|
277
|
+
if cache_config is not None:
|
|
278
|
+
enable_cache_on_transformer(self.pipeline, cache_config)
|
|
279
|
+
|
|
280
|
+
# Configure the schedulers if specified - a pipeline that denoises two
|
|
281
|
+
# modalities against two schedules configures each of them separately
|
|
282
|
+
load_and_configure_scheduler(
|
|
283
|
+
self.pipeline_definition.get("scheduler", None), self.pipeline
|
|
284
|
+
)
|
|
285
|
+
load_and_configure_scheduler(
|
|
286
|
+
self.pipeline_definition.get("audio_scheduler", None),
|
|
287
|
+
self.pipeline,
|
|
288
|
+
"audio_scheduler",
|
|
289
|
+
)
|
|
290
|
+
|
|
291
|
+
self.publish_shared_components(shared_components)
|
|
292
|
+
|
|
293
|
+
# Load and configure LoRA models
|
|
294
|
+
load_loras(self.pipeline_definition.get("loras", []), self.pipeline)
|
|
295
|
+
|
|
296
|
+
# Load and configure IP-Adapter
|
|
297
|
+
load_ip_adapter(self.pipeline_definition.get("ip_adapter", None), self.pipeline)
|
|
298
|
+
|
|
299
|
+
# The adapters are on the pipeline now, so its offloading hooks can be
|
|
300
|
+
# installed over the weights they added
|
|
301
|
+
if adapters_to_load:
|
|
302
|
+
self.pipeline = place_component(
|
|
303
|
+
self.pipeline,
|
|
304
|
+
"pipeline",
|
|
305
|
+
self.configuration,
|
|
306
|
+
self.device,
|
|
307
|
+
# A modular pipeline's manager owns its placement, and the manager
|
|
308
|
+
# load_component gave it is the one it holds
|
|
309
|
+
getattr(self.pipeline, "_components_manager", None),
|
|
310
|
+
)
|
|
311
|
+
|
|
312
|
+
# Place the components the pipeline loaded itself, once everything that alters
|
|
313
|
+
# them - dtypes, adapters, quantized matmuls - has been applied. Offloading hooks
|
|
314
|
+
# installed before those would be fighting them
|
|
315
|
+
configure_components(
|
|
316
|
+
self.pipeline, self.configuration, self.device, reused_components
|
|
317
|
+
)
|
|
318
|
+
|
|
319
|
+
# Set up random generator if needed - no_generator is a boolean, so an
|
|
320
|
+
# explicit false still gets a generator
|
|
321
|
+
if not self.configuration.get("no_generator", False):
|
|
322
|
+
logger.debug("Setting up random generator")
|
|
323
|
+
self.argument_template["generator"] = torch.Generator(
|
|
324
|
+
self.device
|
|
325
|
+
).manual_seed(self.pipeline_definition.get("seed", self.default_seed))
|
|
326
|
+
|
|
327
|
+
# Hand the first run a clean allocator. Loading churns the device even
|
|
328
|
+
# when little of the pipeline stays there - a quantization pass with
|
|
329
|
+
# 'quantization_device' set works on the accelerator and returns the
|
|
330
|
+
# weights to the host, and group offloading moves components off it
|
|
331
|
+
# again - and the cached blocks left behind are the wrong shape for
|
|
332
|
+
# inference. workflow.py does this between steps; a one-step workflow
|
|
333
|
+
# would otherwise run its only step on top of the loading debris
|
|
334
|
+
gc.collect()
|
|
335
|
+
empty_device_cache()
|
|
336
|
+
|
|
337
|
+
logger.debug("Pipeline loaded successfully")
|
|
338
|
+
|
|
339
|
+
def publish_shared_components(self, shared_components):
|
|
340
|
+
"""Store components that will be shared with other pipelines.
|
|
341
|
+
|
|
342
|
+
Called from load, and again by the workflow when a cached pipeline is
|
|
343
|
+
reused - a cache hit skips load entirely, and the shared_components
|
|
344
|
+
dict is fresh every run, so a warm sharing step must republish or a
|
|
345
|
+
later reusing step finds nothing. get_component rather than getattr -
|
|
346
|
+
a modular pipeline registers a component it did not load as None, and
|
|
347
|
+
sharing that None silently would surface as a missing-component error
|
|
348
|
+
inside the step that reused it.
|
|
349
|
+
"""
|
|
350
|
+
for shared_component_name in self.component_names("shared_components"):
|
|
351
|
+
component = get_component(self.pipeline, shared_component_name)
|
|
352
|
+
if component is None:
|
|
353
|
+
raise ValueError(
|
|
354
|
+
f"Cannot share component '{shared_component_name}' - "
|
|
355
|
+
f"{type(self.pipeline).__name__} registers it but has not "
|
|
356
|
+
f"loaded it"
|
|
357
|
+
)
|
|
358
|
+
logger.debug(f"Storing shared component: {shared_component_name}")
|
|
359
|
+
shared_components[shared_component_name] = component
|
|
360
|
+
|
|
361
|
+
@torch.inference_mode()
|
|
362
|
+
def run(self, arguments, previous_pipelines={}):
|
|
363
|
+
"""
|
|
364
|
+
Execute the pipeline with given arguments.
|
|
365
|
+
|
|
366
|
+
Args:
|
|
367
|
+
arguments: Dictionary of arguments for pipeline execution
|
|
368
|
+
previous_pipelines: Dictionary of previously created pipelines
|
|
369
|
+
|
|
370
|
+
Returns:
|
|
371
|
+
Pipeline output or dictionary containing special outputs
|
|
372
|
+
"""
|
|
373
|
+
if self.pipeline is None:
|
|
374
|
+
logger.error("Pipeline not initialized")
|
|
375
|
+
raise ValueError(
|
|
376
|
+
"Pipeline has not been initialized. Call load(device_identifier, shared_components) first."
|
|
377
|
+
)
|
|
378
|
+
|
|
379
|
+
logger.debug(f"Running pipeline with arguments: {arguments}")
|
|
380
|
+
|
|
381
|
+
try:
|
|
382
|
+
# Handle inversion pipeline
|
|
383
|
+
if self.configuration.get("inversion", False):
|
|
384
|
+
logger.debug("Running inversion pipeline")
|
|
385
|
+
invert_arguments = copy.deepcopy(arguments)
|
|
386
|
+
invert_arguments.pop("generator", None)
|
|
387
|
+
inverted_latents, image_latents, latent_image_ids = (
|
|
388
|
+
self.pipeline.invert(**invert_arguments)
|
|
389
|
+
)
|
|
390
|
+
return {
|
|
391
|
+
"inverted_latents": inverted_latents,
|
|
392
|
+
"image_latents": image_latents,
|
|
393
|
+
"latent_image_ids": latent_image_ids,
|
|
394
|
+
}
|
|
395
|
+
|
|
396
|
+
# Handle generation pipeline
|
|
397
|
+
if self.configuration.get("generate", False):
|
|
398
|
+
logger.debug("Running generation pipeline")
|
|
399
|
+
return {"generated_ids": self.pipeline.generate(**arguments)}
|
|
400
|
+
|
|
401
|
+
chain_definition = self.pipeline_definition.get("chain", None)
|
|
402
|
+
if chain_definition is not None:
|
|
403
|
+
from .chain import run_chain
|
|
404
|
+
|
|
405
|
+
logger.debug("Running chained pipeline")
|
|
406
|
+
return run_chain(self, chain_definition, arguments)
|
|
407
|
+
|
|
408
|
+
return self._run_once(arguments)
|
|
409
|
+
|
|
410
|
+
except WorkflowCancelled:
|
|
411
|
+
logger.info("Pipeline run cancelled")
|
|
412
|
+
raise
|
|
413
|
+
except Exception as e:
|
|
414
|
+
# One log line with the full traceback - every error class was
|
|
415
|
+
# logged and re-raised identically
|
|
416
|
+
logger.error(f"{type(e).__name__} running pipeline: {e}", exc_info=True)
|
|
417
|
+
raise
|
|
418
|
+
|
|
419
|
+
def _run_once(self, arguments):
|
|
420
|
+
"""Run one standard pipeline invocation with fully resolved arguments.
|
|
421
|
+
|
|
422
|
+
This is the whole per-call execution path - prompt encoding, the
|
|
423
|
+
pipeline call itself, and output normalization - shared by the single
|
|
424
|
+
run and every segment of a chained run.
|
|
425
|
+
"""
|
|
426
|
+
if self.pipeline_definition.get("remote_text_encoder", None) is not None:
|
|
427
|
+
logger.info("Invoking remote text encoder")
|
|
428
|
+
remote_config = self.pipeline_definition["remote_text_encoder"]
|
|
429
|
+
prompt_embeds = remote_text_encoder(
|
|
430
|
+
arguments.pop("prompt"),
|
|
431
|
+
remote_config.get("url"),
|
|
432
|
+
device=self.device,
|
|
433
|
+
)
|
|
434
|
+
arguments["prompt_embeds"] = prompt_embeds
|
|
435
|
+
elif self.configuration.get("prompt_weighting", False):
|
|
436
|
+
from ..prompt_weighting import apply_prompt_weighting
|
|
437
|
+
|
|
438
|
+
# The step's device override travels with the call - embeddings
|
|
439
|
+
# must land where the transformer runs
|
|
440
|
+
apply_prompt_weighting(self.pipeline, arguments, self.device)
|
|
441
|
+
|
|
442
|
+
# Run standard pipeline
|
|
443
|
+
logger.debug("Running standard pipeline")
|
|
444
|
+
output = self._execute_pipeline(arguments)
|
|
445
|
+
|
|
446
|
+
# A raw tensor result - latents, embeddings - is held for the rest of the
|
|
447
|
+
# workflow, so it rests in system memory instead of occupying the
|
|
448
|
+
# accelerator that the next step needs. Pipelines consuming it place it back
|
|
449
|
+
# on their own device.
|
|
450
|
+
if hasattr(output, "to"):
|
|
451
|
+
logger.debug("Moving tensor output to system memory")
|
|
452
|
+
output = output.to("cpu")
|
|
453
|
+
|
|
454
|
+
attach_audio_sample_rate(self.pipeline, output)
|
|
455
|
+
|
|
456
|
+
return output
|
|
457
|
+
|
|
458
|
+
def _execute_pipeline(self, arguments):
|
|
459
|
+
"""Execute the pipeline with optional TeaCache and attention backend contexts."""
|
|
460
|
+
teacache_config = self.configuration.get("teacache", None)
|
|
461
|
+
attn_backend = self.configuration.get("attention_backend", None)
|
|
462
|
+
|
|
463
|
+
# Determine the execution context
|
|
464
|
+
if teacache_config is not None:
|
|
465
|
+
num_steps = arguments.get("num_inference_steps", None)
|
|
466
|
+
if num_steps is None:
|
|
467
|
+
logger.warning(
|
|
468
|
+
"TeaCache requires num_inference_steps in arguments, running without TeaCache"
|
|
469
|
+
)
|
|
470
|
+
return self._call_pipeline(arguments, attn_backend)
|
|
471
|
+
|
|
472
|
+
rel_l1_thresh = teacache_config.get("rel_l1_thresh", None)
|
|
473
|
+
coefficients = teacache_config.get("coefficients", None)
|
|
474
|
+
variant = teacache_config.get("variant", None)
|
|
475
|
+
with teacache_context(
|
|
476
|
+
self.pipeline, num_steps, rel_l1_thresh, coefficients, variant
|
|
477
|
+
):
|
|
478
|
+
return self._call_pipeline(arguments, attn_backend)
|
|
479
|
+
else:
|
|
480
|
+
return self._call_pipeline(arguments, attn_backend)
|
|
481
|
+
|
|
482
|
+
def _call_pipeline(self, arguments, attn_backend):
|
|
483
|
+
"""Call the pipeline with optional attention backend and cache contexts."""
|
|
484
|
+
arguments = self._with_step_callback(arguments)
|
|
485
|
+
with contextlib.ExitStack() as stack:
|
|
486
|
+
if attn_backend is not None:
|
|
487
|
+
logger.info(f"Using attention backend: {attn_backend}")
|
|
488
|
+
stack.enter_context(attention_backend(attn_backend))
|
|
489
|
+
|
|
490
|
+
stack.enter_context(stateful_cache_context(self.pipeline))
|
|
491
|
+
|
|
492
|
+
return self.pipeline(**arguments)
|
|
493
|
+
|
|
494
|
+
def _with_step_callback(self, arguments):
|
|
495
|
+
"""Inject a callback_on_step_end that reports per-step progress to the
|
|
496
|
+
active run context and raises when the run has been cancelled.
|
|
497
|
+
|
|
498
|
+
Workflow JSON cannot express a callable, so this is the only way a
|
|
499
|
+
diffusion-step callback ever reaches a pipeline call. Only pipelines
|
|
500
|
+
that name the parameter explicitly get one - a **kwargs signature is
|
|
501
|
+
no promise the pipeline honors it.
|
|
502
|
+
"""
|
|
503
|
+
try:
|
|
504
|
+
parameters = inspect.signature(self.pipeline.__call__).parameters
|
|
505
|
+
except (TypeError, ValueError):
|
|
506
|
+
return arguments
|
|
507
|
+
if "callback_on_step_end" not in parameters:
|
|
508
|
+
return arguments
|
|
509
|
+
|
|
510
|
+
run_context = get_context()
|
|
511
|
+
num_steps = arguments.get("num_inference_steps", None)
|
|
512
|
+
|
|
513
|
+
def on_step_end(pipe, step_index, timestep, callback_kwargs):
|
|
514
|
+
run_context.emit(
|
|
515
|
+
"pipeline_step",
|
|
516
|
+
step=step_index + 1,
|
|
517
|
+
total_steps=getattr(pipe, "_num_timesteps", None) or num_steps,
|
|
518
|
+
)
|
|
519
|
+
run_context.check_cancelled()
|
|
520
|
+
return callback_kwargs
|
|
521
|
+
|
|
522
|
+
return {**arguments, "callback_on_step_end": on_step_end}
|
|
523
|
+
|
|
524
|
+
def load_optional_component(
|
|
525
|
+
self, component_name, from_pretrained_arguments, default_device
|
|
526
|
+
):
|
|
527
|
+
"""Load an optional component if specified in pipeline definition."""
|
|
528
|
+
component_definition = self.pipeline_definition.get(component_name, None)
|
|
529
|
+
|
|
530
|
+
if component_definition is not None:
|
|
531
|
+
logger.info(f"Loading component: {component_name}")
|
|
532
|
+
component_configuration = component_definition.get("configuration", None)
|
|
533
|
+
if component_configuration is not None:
|
|
534
|
+
# A copy for the same reason the pipeline's own arguments are copied:
|
|
535
|
+
# what goes in here is consumed by load_component, and the definition
|
|
536
|
+
# is the workflow's, not this load's
|
|
537
|
+
component_from_pretrained_arguments = dict(
|
|
538
|
+
component_definition["from_pretrained_arguments"]
|
|
539
|
+
)
|
|
540
|
+
|
|
541
|
+
# Handle quantization configuration
|
|
542
|
+
quantization_configuration = get_quantization_configuration(
|
|
543
|
+
component_definition
|
|
544
|
+
)
|
|
545
|
+
if quantization_configuration is not None:
|
|
546
|
+
logger.debug(f"Adding quantization config for {component_name}")
|
|
547
|
+
component_from_pretrained_arguments["quantization_config"] = (
|
|
548
|
+
quantization_configuration
|
|
549
|
+
)
|
|
550
|
+
|
|
551
|
+
device = component_configuration.get("device", default_device)
|
|
552
|
+
|
|
553
|
+
component = load_component(
|
|
554
|
+
component_name,
|
|
555
|
+
component_configuration,
|
|
556
|
+
component_from_pretrained_arguments,
|
|
557
|
+
device,
|
|
558
|
+
)
|
|
559
|
+
|
|
560
|
+
logger.debug(f"Loaded optional component: {component_name}")
|
|
561
|
+
from_pretrained_arguments[component_name] = component
|
|
562
|
+
|
|
563
|
+
def configure_loaded_components(self):
|
|
564
|
+
# Configure VAE settings
|
|
565
|
+
vae = self.configuration.get("vae", {})
|
|
566
|
+
if vae.get("enable_slicing", False):
|
|
567
|
+
logger.debug("Enabling VAE slicing")
|
|
568
|
+
self.pipeline.vae.enable_slicing()
|
|
569
|
+
if vae.get("enable_tiling", False):
|
|
570
|
+
logger.debug("Enabling VAE tiling")
|
|
571
|
+
self.pipeline.vae.enable_tiling()
|
|
572
|
+
if vae.get("channels_last", False):
|
|
573
|
+
logger.debug("Setting VAE memory format")
|
|
574
|
+
self.pipeline.vae.to(memory_format=torch.channels_last)
|
|
575
|
+
|
|
576
|
+
# Configure UNet settings
|
|
577
|
+
unet = self.configuration.get("unet", {})
|
|
578
|
+
if unet.get("enable_forward_chunking", False):
|
|
579
|
+
logger.debug("Enabling UNet forward chunking")
|
|
580
|
+
self.pipeline.unet.enable_forward_chunking()
|
|
581
|
+
if unet.get("channels_last", False):
|
|
582
|
+
logger.debug("Setting UNet memory format")
|
|
583
|
+
self.pipeline.unet.to(memory_format=torch.channels_last)
|
|
584
|
+
|
|
585
|
+
# Configure UNet attention processor
|
|
586
|
+
if unet.get("attn_processor_type", None) is not None:
|
|
587
|
+
logger.debug("Enabling UNet custom attention processor")
|
|
588
|
+
attn_processor = unet["attn_processor_type"]()
|
|
589
|
+
self.pipeline.unet.set_attn_processor(attn_processor)
|
|
590
|
+
|
|
591
|
+
# Configure transformer settings
|
|
592
|
+
transformer = self.configuration.get("transformer", {})
|
|
593
|
+
if transformer.get("attn_processor_type", None) is not None:
|
|
594
|
+
logger.debug("Enabling transformer custom attention processor")
|
|
595
|
+
attn_processor = transformer["attn_processor_type"]()
|
|
596
|
+
self.pipeline.transformer.set_attn_processor(attn_processor)
|
|
597
|
+
|
|
598
|
+
# configure optional components
|
|
599
|
+
for component_name in declared_component_names(self.pipeline_definition):
|
|
600
|
+
component_configuration = self.configuration.get(component_name, None)
|
|
601
|
+
if component_configuration is None:
|
|
602
|
+
continue
|
|
603
|
+
|
|
604
|
+
# get_component() raises on a genuinely missing attribute (a typo) and
|
|
605
|
+
# returns None for one that is registered but unloaded - both cases are
|
|
606
|
+
# unconfigurable, so both are skipped here exactly as a plain missing
|
|
607
|
+
# component always was
|
|
608
|
+
try:
|
|
609
|
+
component = get_component(self.pipeline, component_name)
|
|
610
|
+
except ValueError:
|
|
611
|
+
component = None
|
|
612
|
+
|
|
613
|
+
if component is not None:
|
|
614
|
+
logger.debug(f"Configuring optional component: {component_name}")
|
|
615
|
+
torch_dtype = component_configuration.get("torch_dtype", None)
|
|
616
|
+
if torch_dtype is not None:
|
|
617
|
+
logger.debug(f"Setting {component_name} torch dtype: {torch_dtype}")
|
|
618
|
+
component.to(torch_dtype)
|
|
619
|
+
|
|
620
|
+
|
|
621
|
+
def configure_components(pipeline, configuration, default_device, reused_components=()):
|
|
622
|
+
"""Place the components a pipeline loaded for itself.
|
|
623
|
+
|
|
624
|
+
A modular pipeline pulls its own component weights, so they are only reachable once
|
|
625
|
+
the pipeline is loaded - too late for the offloading load_component sets up. Group
|
|
626
|
+
offloading a component here streams it between system memory and the accelerator a
|
|
627
|
+
piece at a time, which is what fits a pipeline whose components are each larger than
|
|
628
|
+
the device.
|
|
629
|
+
|
|
630
|
+
A component this step reused is not one it loaded: it already carries the placement
|
|
631
|
+
the step that shared it gave it, and offloading hooks do not survive being applied
|
|
632
|
+
twice. Those are skipped, so a workflow can reuse a component into a step whose
|
|
633
|
+
configuration was written for loading it.
|
|
634
|
+
|
|
635
|
+
Args:
|
|
636
|
+
pipeline: The loaded pipeline
|
|
637
|
+
configuration: Pipeline configuration dictionary
|
|
638
|
+
default_device: Device the pipeline runs on
|
|
639
|
+
reused_components: Names of the components an earlier step shared into this one
|
|
640
|
+
"""
|
|
641
|
+
for component_name, component_configuration in configuration.get(
|
|
642
|
+
"components", {}
|
|
643
|
+
).items():
|
|
644
|
+
# A dotted path reaches inside a component, and it is the component itself
|
|
645
|
+
# that was shared - 'text_encoder.model' belongs to a reused 'text_encoder'
|
|
646
|
+
if component_name.split(".")[0] in reused_components:
|
|
647
|
+
logger.info(
|
|
648
|
+
f"Component '{component_name}' was shared by an earlier step - "
|
|
649
|
+
"keeping the placement that step gave it"
|
|
650
|
+
)
|
|
651
|
+
continue
|
|
652
|
+
|
|
653
|
+
component = get_component(pipeline, component_name)
|
|
654
|
+
if component is None:
|
|
655
|
+
# Registered but unloaded (e.g. a components map reused across workflow
|
|
656
|
+
# selections, or a component diffusers warned-and-skipped past at load) -
|
|
657
|
+
# skip just this entry rather than aborting the whole run
|
|
658
|
+
logger.warning(
|
|
659
|
+
f"Component '{component_name}' is not loaded (workflow selection "
|
|
660
|
+
"may not use it) - skipping its configuration"
|
|
661
|
+
)
|
|
662
|
+
continue
|
|
663
|
+
|
|
664
|
+
# Prune before any hooks are installed - an offload hook pins and tracks
|
|
665
|
+
# exactly the modules that exist when it is applied, so pruning afterwards
|
|
666
|
+
# would leave it streaming weights that can never run
|
|
667
|
+
truncate_module_lists(component, component_name, component_configuration)
|
|
668
|
+
replace_modules_with_identity(
|
|
669
|
+
component, component_name, component_configuration
|
|
670
|
+
)
|
|
671
|
+
|
|
672
|
+
group_offload_configuration = get_group_offload_configuration(
|
|
673
|
+
component_configuration, default_device
|
|
674
|
+
)
|
|
675
|
+
if group_offload_configuration is not None:
|
|
676
|
+
# apply_group_offloading rather than the component's own
|
|
677
|
+
# enable_group_offload - a component may be a transformers model, or a
|
|
678
|
+
# module inside one, and only diffusers models have the method
|
|
679
|
+
from diffusers.hooks import apply_group_offloading
|
|
680
|
+
|
|
681
|
+
logger.info(f"Group offloading {component_name}")
|
|
682
|
+
apply_group_offloading(component, **group_offload_configuration)
|
|
683
|
+
|
|
684
|
+
# Tiled decoding, for a component that decodes but is not the one called
|
|
685
|
+
# 'vae' - LTX-2.5's diffusion decoder, which decodes the whole video volume
|
|
686
|
+
# in one allocation unless it is told to tile
|
|
687
|
+
enable_tiling(component, component_name, component_configuration)
|
|
688
|
+
|
|
689
|
+
device = component_configuration.get("device", None)
|
|
690
|
+
residency = component_configuration.get("residency", "resident")
|
|
691
|
+
if residency == "on_demand":
|
|
692
|
+
apply_on_demand_placement(
|
|
693
|
+
component,
|
|
694
|
+
component_name,
|
|
695
|
+
device if device is not None else default_device,
|
|
696
|
+
group_offload_configuration is not None,
|
|
697
|
+
)
|
|
698
|
+
elif device is not None:
|
|
699
|
+
logger.info(f"Moving {component_name} to device: {device}")
|
|
700
|
+
component.to(device)
|
|
701
|
+
|
|
702
|
+
# A compiled component should pin its attention backend - the per-call
|
|
703
|
+
# attention_backend context manager would switch implementations under a
|
|
704
|
+
# compiled graph and force a recompile on every run
|
|
705
|
+
component_attention_backend = component_configuration.get(
|
|
706
|
+
"attention_backend", None
|
|
707
|
+
)
|
|
708
|
+
if component_attention_backend is not None:
|
|
709
|
+
logger.info(
|
|
710
|
+
f"Setting {component_name} attention backend: {component_attention_backend}"
|
|
711
|
+
)
|
|
712
|
+
component.set_attention_backend(component_attention_backend)
|
|
713
|
+
|
|
714
|
+
# Compile last - the graph must capture final dtypes, adapters,
|
|
715
|
+
# quantization, and offload hooks
|
|
716
|
+
compile_configuration = component_configuration.get("compile", None)
|
|
717
|
+
if compile_configuration is not None:
|
|
718
|
+
apply_compile(
|
|
719
|
+
component,
|
|
720
|
+
component_name,
|
|
721
|
+
compile_configuration,
|
|
722
|
+
device if device is not None else default_device,
|
|
723
|
+
)
|
|
724
|
+
|
|
725
|
+
|
|
726
|
+
def enable_tiling(component, component_name, component_configuration):
|
|
727
|
+
"""Turn on tiled decoding for a component whose configuration asks for it.
|
|
728
|
+
|
|
729
|
+
The pipeline-level `vae` block covers the component actually named 'vae'. This
|
|
730
|
+
covers any other component that decodes - LTX-2.5's `diffusion_decoder`, which
|
|
731
|
+
otherwise decodes the whole video volume in a single allocation and asks for
|
|
732
|
+
tens of GiB at 2x resolutions. `true` takes the model's own default tile size;
|
|
733
|
+
a dict passes the tile and stride sizes through, which is what a card smaller
|
|
734
|
+
than those defaults needs.
|
|
735
|
+
|
|
736
|
+
Args:
|
|
737
|
+
component: The loaded component
|
|
738
|
+
component_name: Its name, for logging and errors
|
|
739
|
+
component_configuration: That component's configuration block
|
|
740
|
+
|
|
741
|
+
Raises:
|
|
742
|
+
ValueError: If the component has no enable_tiling() to call
|
|
743
|
+
"""
|
|
744
|
+
tiling = component_configuration.get("enable_tiling", False)
|
|
745
|
+
if not tiling:
|
|
746
|
+
return
|
|
747
|
+
|
|
748
|
+
if not has_method(component, "enable_tiling"):
|
|
749
|
+
raise ValueError(
|
|
750
|
+
f"'{component_name}' does not support tiling - "
|
|
751
|
+
f"{type(component).__name__} has no enable_tiling()"
|
|
752
|
+
)
|
|
753
|
+
|
|
754
|
+
arguments = tiling if isinstance(tiling, dict) else {}
|
|
755
|
+
logger.info(
|
|
756
|
+
f"Enabling tiling on {component_name}"
|
|
757
|
+
+ (f" with {', '.join(arguments)}" if arguments else "")
|
|
758
|
+
)
|
|
759
|
+
component.enable_tiling(**arguments)
|
|
760
|
+
|
|
761
|
+
|
|
762
|
+
def _resolve_submodule(component, component_name, path):
|
|
763
|
+
"""Follow a dotted path from a component to a module inside it.
|
|
764
|
+
|
|
765
|
+
Args:
|
|
766
|
+
component: The component the path starts from
|
|
767
|
+
component_name: Its name, for errors
|
|
768
|
+
path: Dotted attribute path relative to the component, e.g.
|
|
769
|
+
'language_model.layers'
|
|
770
|
+
|
|
771
|
+
Returns:
|
|
772
|
+
(parent, attribute_name, module) - the module and where it hangs
|
|
773
|
+
|
|
774
|
+
Raises:
|
|
775
|
+
ValueError: If any step of the path is not an attribute
|
|
776
|
+
"""
|
|
777
|
+
parent = None
|
|
778
|
+
module = component
|
|
779
|
+
for attribute_name in path.split("."):
|
|
780
|
+
parent = module
|
|
781
|
+
module = getattr(parent, attribute_name, _MISSING)
|
|
782
|
+
if module is _MISSING:
|
|
783
|
+
raise ValueError(
|
|
784
|
+
f"'{component_name}' has no module at '{path}' - "
|
|
785
|
+
f"{type(parent).__name__} has no attribute '{attribute_name}'"
|
|
786
|
+
)
|
|
787
|
+
return parent, path.rsplit(".", 1)[-1], module
|
|
788
|
+
|
|
789
|
+
|
|
790
|
+
def truncate_module_lists(component, component_name, component_configuration):
|
|
791
|
+
"""Drop the tail of a ModuleList a run never reads.
|
|
792
|
+
|
|
793
|
+
An encoder used for its hidden states can run layers whose output nothing
|
|
794
|
+
consumes: MiniMax-H3 conditions on hidden_states[50] of its 64-layer
|
|
795
|
+
Qwen3-VL, so layers 51-63 compute - and, offloaded, stream from system
|
|
796
|
+
memory - for nothing, on every encode. Keeping 51 layers leaves
|
|
797
|
+
hidden_states[50] bit-identical (index 50 of the returned tuple is the
|
|
798
|
+
input to layer 50, recorded before it runs; keeping only 50 would make it
|
|
799
|
+
the final-norm output instead, which is a different tensor).
|
|
800
|
+
|
|
801
|
+
The configuration maps a dotted path inside the component to the number of
|
|
802
|
+
entries to keep:
|
|
803
|
+
|
|
804
|
+
"truncate_layers": { "language_model.layers": 51 }
|
|
805
|
+
|
|
806
|
+
Truncation is in place, so the component's registration on its pipeline and
|
|
807
|
+
its config are untouched - a block that validates against
|
|
808
|
+
config.num_hidden_layers still sees the checkpoint's own count.
|
|
809
|
+
|
|
810
|
+
Args:
|
|
811
|
+
component: The loaded component
|
|
812
|
+
component_name: Its name, for logging and errors
|
|
813
|
+
component_configuration: That component's configuration block
|
|
814
|
+
|
|
815
|
+
Raises:
|
|
816
|
+
ValueError: If a path does not lead to a ModuleList, or keep is not
|
|
817
|
+
a positive count
|
|
818
|
+
"""
|
|
819
|
+
truncations = component_configuration.get("truncate_layers", None)
|
|
820
|
+
if not truncations:
|
|
821
|
+
return
|
|
822
|
+
|
|
823
|
+
for path, keep in truncations.items():
|
|
824
|
+
_, _, module_list = _resolve_submodule(component, component_name, path)
|
|
825
|
+
if not isinstance(module_list, torch.nn.ModuleList):
|
|
826
|
+
raise ValueError(
|
|
827
|
+
f"'{component_name}' cannot truncate '{path}' - it is a "
|
|
828
|
+
f"{type(module_list).__name__}, not a ModuleList"
|
|
829
|
+
)
|
|
830
|
+
keep = int(keep)
|
|
831
|
+
if keep < 1:
|
|
832
|
+
raise ValueError(
|
|
833
|
+
f"'{component_name}' truncate_layers keeps {keep} of '{path}' - "
|
|
834
|
+
"at least one layer has to remain"
|
|
835
|
+
)
|
|
836
|
+
if keep >= len(module_list):
|
|
837
|
+
logger.warning(
|
|
838
|
+
f"'{component_name}' truncate_layers keeps {keep} of '{path}', "
|
|
839
|
+
f"which already has {len(module_list)} - nothing to drop"
|
|
840
|
+
)
|
|
841
|
+
continue
|
|
842
|
+
logger.info(
|
|
843
|
+
f"Truncating {component_name} '{path}' from {len(module_list)} "
|
|
844
|
+
f"layers to {keep}"
|
|
845
|
+
)
|
|
846
|
+
del module_list[keep:]
|
|
847
|
+
|
|
848
|
+
|
|
849
|
+
def replace_modules_with_identity(component, component_name, component_configuration):
|
|
850
|
+
"""Swap out modules a run never calls, freeing what they hold.
|
|
851
|
+
|
|
852
|
+
For weights that exist on the checkpoint but are outside the path the
|
|
853
|
+
pipeline actually runs - a language-model head on a model used as an
|
|
854
|
+
encoder, say. The module is replaced with an Identity so the model keeps
|
|
855
|
+
its shape for anything that looks the attribute up, while its parameters
|
|
856
|
+
are dropped rather than held (and, offloaded, pinned) for a call that
|
|
857
|
+
never comes.
|
|
858
|
+
|
|
859
|
+
"remove_modules": [ "lm_head" ]
|
|
860
|
+
|
|
861
|
+
Args:
|
|
862
|
+
component: The loaded component
|
|
863
|
+
component_name: Its name, for logging and errors
|
|
864
|
+
component_configuration: That component's configuration block
|
|
865
|
+
|
|
866
|
+
Raises:
|
|
867
|
+
ValueError: If a named module does not exist on the component
|
|
868
|
+
"""
|
|
869
|
+
for path in component_configuration.get("remove_modules", []):
|
|
870
|
+
parent, attribute_name, _ = _resolve_submodule(component, component_name, path)
|
|
871
|
+
logger.info(f"Replacing {component_name} '{path}' with Identity")
|
|
872
|
+
setattr(parent, attribute_name, torch.nn.Identity())
|
|
873
|
+
|
|
874
|
+
|
|
875
|
+
# The calls that mean "this component is working now". A component is moved to
|
|
876
|
+
# the accelerator around whichever of these it actually defines
|
|
877
|
+
_ON_DEMAND_ENTRY_POINTS = ("forward", "encode", "decode")
|
|
878
|
+
|
|
879
|
+
|
|
880
|
+
def apply_on_demand_placement(
|
|
881
|
+
component, component_name, device, group_offloaded, offload_device="cpu"
|
|
882
|
+
):
|
|
883
|
+
"""Keep a component in system memory and move it to the device only while it runs.
|
|
884
|
+
|
|
885
|
+
Sits between the two placements dw already has. A 'device' component is resident
|
|
886
|
+
for the whole run, which wastes the accelerator on something used twice; group
|
|
887
|
+
offloading streams per submodule forward, which restreams the whole model once
|
|
888
|
+
per call of every leaf - ruinous for a VAE, whose tiled decode calls its blocks
|
|
889
|
+
once per tile. This moves the model as a whole around each entry point, so a
|
|
890
|
+
tiling loop sits inside a single pair of transfers.
|
|
891
|
+
|
|
892
|
+
That trade only pays for components called a handful of times per run. A
|
|
893
|
+
denoising transformer is called once per step, so per-call transfers would cost
|
|
894
|
+
far more than they save - group offloading is the tool for those.
|
|
895
|
+
|
|
896
|
+
Args:
|
|
897
|
+
component: The component to place
|
|
898
|
+
component_name: Name of the component, for logging
|
|
899
|
+
device: Device to run the component on
|
|
900
|
+
group_offloaded: Whether group offloading was applied to this component
|
|
901
|
+
offload_device: Where the component rests between calls
|
|
902
|
+
|
|
903
|
+
Raises:
|
|
904
|
+
ValueError: If the component is also group offloaded
|
|
905
|
+
"""
|
|
906
|
+
if group_offloaded:
|
|
907
|
+
raise ValueError(
|
|
908
|
+
f"Component '{component_name}' sets both 'group_offload' and "
|
|
909
|
+
"'residency: on_demand'. A group offloaded module holds one group at a "
|
|
910
|
+
"time and ignores the whole-model moves on-demand placement makes, so "
|
|
911
|
+
"the two cannot both own its placement - pick one"
|
|
912
|
+
)
|
|
913
|
+
|
|
914
|
+
if get_device_type(device) == "cpu":
|
|
915
|
+
# Nothing to move it off of, so the wrappers would be pure overhead
|
|
916
|
+
logger.debug(
|
|
917
|
+
f"Ignoring 'residency: on_demand' for {component_name} - {device} is the "
|
|
918
|
+
"device it would rest on anyway"
|
|
919
|
+
)
|
|
920
|
+
return
|
|
921
|
+
|
|
922
|
+
component.to(offload_device)
|
|
923
|
+
|
|
924
|
+
# One depth counter for the whole component, not one per entry point: decode()
|
|
925
|
+
# calls forward() internally, and an inner return must not offload the model
|
|
926
|
+
# out from under the call that is still running
|
|
927
|
+
state = {"depth": 0}
|
|
928
|
+
|
|
929
|
+
def wrap(entry_point):
|
|
930
|
+
original = getattr(component, entry_point, None)
|
|
931
|
+
if not callable(original):
|
|
932
|
+
return False
|
|
933
|
+
|
|
934
|
+
@functools.wraps(original)
|
|
935
|
+
def on_demand(*args, **kwargs):
|
|
936
|
+
if state["depth"] == 0:
|
|
937
|
+
component.to(device)
|
|
938
|
+
state["depth"] += 1
|
|
939
|
+
try:
|
|
940
|
+
return original(*args, **kwargs)
|
|
941
|
+
finally:
|
|
942
|
+
state["depth"] -= 1
|
|
943
|
+
if state["depth"] == 0:
|
|
944
|
+
component.to(offload_device)
|
|
945
|
+
# Hand the freed space back to the driver rather than leaving it
|
|
946
|
+
# reserved - the headroom is the entire point of doing this
|
|
947
|
+
empty_device_cache()
|
|
948
|
+
|
|
949
|
+
# functools.wraps carries __wrapped__, so inspect.signature() still reports
|
|
950
|
+
# the real parameters. Callers introspect them: MiniMax H3's denoiser picks
|
|
951
|
+
# which arguments to pass by reading signature(transformer.forward)
|
|
952
|
+
setattr(component, entry_point, on_demand)
|
|
953
|
+
return True
|
|
954
|
+
|
|
955
|
+
wrapped = [name for name in _ON_DEMAND_ENTRY_POINTS if wrap(name)]
|
|
956
|
+
if not wrapped:
|
|
957
|
+
raise ValueError(
|
|
958
|
+
f"Component '{component_name}' sets 'residency: on_demand' but defines "
|
|
959
|
+
f"none of {', '.join(_ON_DEMAND_ENTRY_POINTS)}, so there is no call to "
|
|
960
|
+
"move it around"
|
|
961
|
+
)
|
|
962
|
+
logger.info(
|
|
963
|
+
f"Placing {component_name} on demand: resting on {offload_device}, "
|
|
964
|
+
f"running on {device} around {', '.join(wrapped)}"
|
|
965
|
+
)
|
|
966
|
+
|
|
967
|
+
|
|
968
|
+
def apply_compile(component, component_name, compile_configuration, device):
|
|
969
|
+
"""Compile a component with torch.compile.
|
|
970
|
+
|
|
971
|
+
Compilation happens in place (nn.Module.compile) so the module stays registered
|
|
972
|
+
on its pipeline. With 'repeated_blocks' true, only the model's repeated block
|
|
973
|
+
classes are compiled (diffusers' regional compilation) - near the same speedup
|
|
974
|
+
as full compilation with a fraction of the cold-start cost.
|
|
975
|
+
|
|
976
|
+
Args:
|
|
977
|
+
component: The component to compile
|
|
978
|
+
component_name: Name of the component, for logging
|
|
979
|
+
compile_configuration: Dict of options - 'repeated_blocks' selects regional
|
|
980
|
+
compilation, everything else ('mode', 'fullgraph', 'dynamic', ...) is
|
|
981
|
+
passed to torch.compile
|
|
982
|
+
device: Device the component runs on
|
|
983
|
+
"""
|
|
984
|
+
# Inductor support on MPS is too immature to be worth the compile time
|
|
985
|
+
if get_device_type(device) == "mps":
|
|
986
|
+
logger.warning(
|
|
987
|
+
f"torch.compile is not supported on MPS, skipping {component_name}"
|
|
988
|
+
)
|
|
989
|
+
return
|
|
990
|
+
|
|
991
|
+
options = dict(compile_configuration)
|
|
992
|
+
repeated_blocks = options.pop("repeated_blocks", False)
|
|
993
|
+
|
|
994
|
+
if repeated_blocks:
|
|
995
|
+
if not has_method(component, "compile_repeated_blocks"):
|
|
996
|
+
raise ValueError(
|
|
997
|
+
f"repeated_blocks compilation requires a diffusers model with "
|
|
998
|
+
f"repeated block support, {type(component).__name__} does not have it"
|
|
999
|
+
)
|
|
1000
|
+
logger.info(f"Compiling repeated blocks of {component_name}")
|
|
1001
|
+
component.compile_repeated_blocks(**options)
|
|
1002
|
+
else:
|
|
1003
|
+
logger.info(f"Compiling {component_name}")
|
|
1004
|
+
component.compile(**options)
|
|
1005
|
+
|
|
1006
|
+
|
|
1007
|
+
_MISSING = object()
|
|
1008
|
+
|
|
1009
|
+
|
|
1010
|
+
def get_component(pipeline, component_name):
|
|
1011
|
+
"""Look a component up on a pipeline, by name or by a dotted path into it.
|
|
1012
|
+
|
|
1013
|
+
A dotted path reaches a module inside a component, which is how a component that
|
|
1014
|
+
holds the model rather than being one - a transformers model wrapping its own - is
|
|
1015
|
+
offloaded.
|
|
1016
|
+
|
|
1017
|
+
A modular pipeline registers a component it has not loaded (a workflow selection
|
|
1018
|
+
that does not use it, or one diffusers warned-and-skipped past) as a None-valued
|
|
1019
|
+
attribute rather than omitting it entirely - that is a real attribute, not a typo,
|
|
1020
|
+
so it is returned as None rather than raising. A missing attribute is still a hard
|
|
1021
|
+
error: it means the name itself is wrong. Callers decide what "unloaded" should
|
|
1022
|
+
mean for them (skip with a warning, skip silently, ...); this just tells them apart.
|
|
1023
|
+
|
|
1024
|
+
Args:
|
|
1025
|
+
pipeline: The loaded pipeline
|
|
1026
|
+
component_name: Name of the component, e.g. 'vae' or 'text_encoder.model'
|
|
1027
|
+
|
|
1028
|
+
Returns:
|
|
1029
|
+
The named component, or None if it (or a step along a dotted path) is
|
|
1030
|
+
registered but not loaded
|
|
1031
|
+
|
|
1032
|
+
Raises:
|
|
1033
|
+
ValueError: If the pipeline has no attribute by that name (or dotted path)
|
|
1034
|
+
"""
|
|
1035
|
+
component = pipeline
|
|
1036
|
+
for attribute_name in component_name.split("."):
|
|
1037
|
+
component = getattr(component, attribute_name, _MISSING)
|
|
1038
|
+
if component is _MISSING:
|
|
1039
|
+
raise ValueError(
|
|
1040
|
+
f"{type(pipeline).__name__} has no component '{component_name}'"
|
|
1041
|
+
)
|
|
1042
|
+
if component is None:
|
|
1043
|
+
return None
|
|
1044
|
+
|
|
1045
|
+
return component
|
|
1046
|
+
|
|
1047
|
+
|
|
1048
|
+
def attach_audio_sample_rate(pipeline, output):
|
|
1049
|
+
"""Record the vocoder's sample rate on an output that carries generated audio.
|
|
1050
|
+
|
|
1051
|
+
Pipelines that generate audio with their video (LTX-2) return the waveform without
|
|
1052
|
+
its sample rate - only the vocoder that produced it knows that. Saving the video
|
|
1053
|
+
needs the rate to mux the audio, so it travels with the output.
|
|
1054
|
+
|
|
1055
|
+
Args:
|
|
1056
|
+
pipeline: The pipeline that produced the output
|
|
1057
|
+
output: The pipeline output
|
|
1058
|
+
"""
|
|
1059
|
+
if getattr(output, "audio", None) is None:
|
|
1060
|
+
return
|
|
1061
|
+
|
|
1062
|
+
vocoder_config = getattr(getattr(pipeline, "vocoder", None), "config", None)
|
|
1063
|
+
sample_rate = getattr(vocoder_config, "output_sampling_rate", None)
|
|
1064
|
+
if sample_rate is None:
|
|
1065
|
+
logger.warning(
|
|
1066
|
+
"Pipeline generated audio but has no vocoder sample rate - "
|
|
1067
|
+
"set 'audio_sample_rate' in the step result to save it with the video"
|
|
1068
|
+
)
|
|
1069
|
+
return
|
|
1070
|
+
|
|
1071
|
+
logger.debug(f"Generated audio has a sample rate of {sample_rate}Hz")
|
|
1072
|
+
output.audio_sample_rate = sample_rate
|
|
1073
|
+
|
|
1074
|
+
|
|
1075
|
+
def load_loras(loras, pipeline):
|
|
1076
|
+
"""Load and configure LoRA models."""
|
|
1077
|
+
adapter_names = []
|
|
1078
|
+
adapter_weights = []
|
|
1079
|
+
|
|
1080
|
+
for i, lora in enumerate(loras):
|
|
1081
|
+
model_name = lora.pop("model_name", None)
|
|
1082
|
+
logger.info(f"Loading LoRA: {model_name}")
|
|
1083
|
+
|
|
1084
|
+
# Use provided adapter_name or generate from index
|
|
1085
|
+
adapter_name = lora.pop("adapter_name", str(i))
|
|
1086
|
+
adapter_names.append(adapter_name)
|
|
1087
|
+
|
|
1088
|
+
# Extract scale for adapter weights - float() because the schema takes a
|
|
1089
|
+
# 'variable:' reference here, and a variable declared as a string default
|
|
1090
|
+
# substitutes as one
|
|
1091
|
+
scale = float(lora.pop("scale", 1.0))
|
|
1092
|
+
adapter_weights.append(scale)
|
|
1093
|
+
|
|
1094
|
+
# Load the LoRA with the adapter name
|
|
1095
|
+
pipeline.load_lora_weights(model_name, adapter_name=adapter_name, **lora)
|
|
1096
|
+
|
|
1097
|
+
# Set adapter weights for all loaded LoRAs
|
|
1098
|
+
if adapter_names:
|
|
1099
|
+
logger.info(
|
|
1100
|
+
f"Setting adapter weights: {list(zip(adapter_names, adapter_weights))}"
|
|
1101
|
+
)
|
|
1102
|
+
# Positionally - diffusers' mixin calls the second parameter 'adapter_weights'
|
|
1103
|
+
# while custom pipelines that delegate to the model (ostris/Krea2OstrisEdit)
|
|
1104
|
+
# call it 'weights'
|
|
1105
|
+
pipeline.set_adapters(adapter_names, adapter_weights)
|
|
1106
|
+
|
|
1107
|
+
|
|
1108
|
+
def load_ip_adapter(ip_adapter_definition, pipeline):
|
|
1109
|
+
"""Load and configure IP-Adapter if specified."""
|
|
1110
|
+
if ip_adapter_definition is not None:
|
|
1111
|
+
model_name = ip_adapter_definition.pop("model_name")
|
|
1112
|
+
logger.info(f"Loading IP-Adapter: {model_name}")
|
|
1113
|
+
scale = ip_adapter_definition.pop("scale", None)
|
|
1114
|
+
pipeline.load_ip_adapter(model_name, **ip_adapter_definition)
|
|
1115
|
+
if scale is not None:
|
|
1116
|
+
pipeline.set_ip_adapter_scale(scale)
|
|
1117
|
+
|
|
1118
|
+
|
|
1119
|
+
def load_and_configure_scheduler(
|
|
1120
|
+
scheduler_definition, pipeline, component_name="scheduler"
|
|
1121
|
+
):
|
|
1122
|
+
"""Load and configure a pipeline's scheduler if specified.
|
|
1123
|
+
|
|
1124
|
+
A definition does either or both of two things, in that order: replace the
|
|
1125
|
+
scheduler with one built from another type's config, and set the sigma
|
|
1126
|
+
shift on whatever scheduler the pipeline then holds.
|
|
1127
|
+
|
|
1128
|
+
The component is named rather than assumed because a pipeline can carry
|
|
1129
|
+
more than one. MiniMax-H3 steps video and audio latents down two schedules
|
|
1130
|
+
inside a single transformer call - 'scheduler' and 'audio_scheduler', whose
|
|
1131
|
+
shifts (12.0 and 3.0 in the released checkpoint) are set independently, and
|
|
1132
|
+
the video one is what a few-step schedule has to lower: at the checkpoint's
|
|
1133
|
+
12.0 a five-point sigma grid spends every step above 0.8 and then drops to
|
|
1134
|
+
zero in one, which denoises to noise.
|
|
1135
|
+
|
|
1136
|
+
Args:
|
|
1137
|
+
scheduler_definition: The step's scheduler block, or None
|
|
1138
|
+
pipeline: The loaded pipeline
|
|
1139
|
+
component_name: Which scheduler the definition configures
|
|
1140
|
+
"""
|
|
1141
|
+
if scheduler_definition is None:
|
|
1142
|
+
return
|
|
1143
|
+
|
|
1144
|
+
scheduler_configuration = scheduler_definition.get("configuration", None) or {}
|
|
1145
|
+
scheduler_type = scheduler_configuration.get("scheduler_type", None)
|
|
1146
|
+
if scheduler_type is not None:
|
|
1147
|
+
from_config_args = scheduler_definition.get("from_config_args", {})
|
|
1148
|
+
logger.info(f"Loading {component_name}: {scheduler_type}")
|
|
1149
|
+
setattr(
|
|
1150
|
+
pipeline,
|
|
1151
|
+
component_name,
|
|
1152
|
+
scheduler_type.from_config(
|
|
1153
|
+
get_component(pipeline, component_name).config, **from_config_args
|
|
1154
|
+
),
|
|
1155
|
+
)
|
|
1156
|
+
|
|
1157
|
+
shift = scheduler_definition.get("shift", None)
|
|
1158
|
+
if shift is None:
|
|
1159
|
+
return
|
|
1160
|
+
|
|
1161
|
+
scheduler = get_component(pipeline, component_name)
|
|
1162
|
+
if scheduler is None:
|
|
1163
|
+
raise ValueError(
|
|
1164
|
+
f"Cannot set a shift on '{component_name}' - the pipeline registers "
|
|
1165
|
+
"it but has not loaded it"
|
|
1166
|
+
)
|
|
1167
|
+
if not has_method(scheduler, "set_shift"):
|
|
1168
|
+
raise ValueError(
|
|
1169
|
+
f"{type(scheduler).__name__} does not take a sigma shift - "
|
|
1170
|
+
f"'{component_name}' has no set_shift()"
|
|
1171
|
+
)
|
|
1172
|
+
|
|
1173
|
+
# Instance state the scheduler keeps until its next set_timesteps, which is
|
|
1174
|
+
# the run itself - so this survives loading and every later run of the step
|
|
1175
|
+
logger.info(f"Setting {component_name} shift: {shift}")
|
|
1176
|
+
scheduler.set_shift(float(shift))
|
|
1177
|
+
|
|
1178
|
+
|
|
1179
|
+
def auto_cpu_offload_enabled(configuration):
|
|
1180
|
+
"""Whether the configuration asks its components manager to offload to the CPU."""
|
|
1181
|
+
return configuration.get("components_manager", {}).get(
|
|
1182
|
+
"enable_auto_cpu_offload", False
|
|
1183
|
+
)
|
|
1184
|
+
|
|
1185
|
+
|
|
1186
|
+
def auto_cpu_offload_active(configuration, device):
|
|
1187
|
+
"""Whether the components manager actually owns device placement.
|
|
1188
|
+
|
|
1189
|
+
Mirrors the MPS skip in create_components_manager() - on MPS the manager
|
|
1190
|
+
never installs its offload hooks, so callers must not assume it owns
|
|
1191
|
+
device placement there.
|
|
1192
|
+
"""
|
|
1193
|
+
return auto_cpu_offload_enabled(configuration) and get_device_type(device) != "mps"
|
|
1194
|
+
|
|
1195
|
+
|
|
1196
|
+
def create_components_manager(configuration, device):
|
|
1197
|
+
"""Create the components manager for a modular pipeline, when one is configured.
|
|
1198
|
+
|
|
1199
|
+
A ComponentsManager tracks the components of a modular pipeline and can keep only
|
|
1200
|
+
the ones currently running on the device, moving the rest to system memory.
|
|
1201
|
+
|
|
1202
|
+
Args:
|
|
1203
|
+
configuration: Pipeline configuration dictionary
|
|
1204
|
+
device: Device the pipeline runs on
|
|
1205
|
+
|
|
1206
|
+
Returns:
|
|
1207
|
+
A configured ComponentsManager, or None when the pipeline does not use one
|
|
1208
|
+
"""
|
|
1209
|
+
manager_configuration = configuration.get("components_manager", None)
|
|
1210
|
+
if manager_configuration is None:
|
|
1211
|
+
return None
|
|
1212
|
+
|
|
1213
|
+
# Imported here because importing modular diffusers warns that it is experimental
|
|
1214
|
+
from diffusers import ComponentsManager
|
|
1215
|
+
|
|
1216
|
+
logger.info("Creating components manager")
|
|
1217
|
+
components_manager = ComponentsManager()
|
|
1218
|
+
|
|
1219
|
+
if auto_cpu_offload_enabled(configuration):
|
|
1220
|
+
# ComponentsManager.enable_auto_cpu_offload() calls device.mem_get_info(),
|
|
1221
|
+
# which torch does not implement for MPS. Unified memory also makes the
|
|
1222
|
+
# feature far less useful there than on CUDA, so skip it rather than fail.
|
|
1223
|
+
if get_device_type(device) == "mps":
|
|
1224
|
+
logger.warning(
|
|
1225
|
+
"components_manager auto CPU offload is not supported on MPS, skipping"
|
|
1226
|
+
)
|
|
1227
|
+
else:
|
|
1228
|
+
offload_arguments = {}
|
|
1229
|
+
memory_reserve_margin = manager_configuration.get(
|
|
1230
|
+
"memory_reserve_margin", None
|
|
1231
|
+
)
|
|
1232
|
+
if memory_reserve_margin is not None:
|
|
1233
|
+
offload_arguments["memory_reserve_margin"] = memory_reserve_margin
|
|
1234
|
+
|
|
1235
|
+
# Enabled before the components load so each one is hooked as it is added
|
|
1236
|
+
logger.info(f"Enabling components manager auto CPU offload on {device}")
|
|
1237
|
+
components_manager.enable_auto_cpu_offload(
|
|
1238
|
+
device=device, **offload_arguments
|
|
1239
|
+
)
|
|
1240
|
+
|
|
1241
|
+
return components_manager
|
|
1242
|
+
|
|
1243
|
+
|
|
1244
|
+
def has_component_group_offload(configuration):
|
|
1245
|
+
"""Whether a per-component entry keeps its component off the device.
|
|
1246
|
+
|
|
1247
|
+
The 'components' block is applied by configure_components() after the pipeline is
|
|
1248
|
+
loaded, but load_component() has to decide where to materialize weights and whether
|
|
1249
|
+
to move the pipeline to the device before that block is ever read. A workflow whose
|
|
1250
|
+
only offload configuration lives under components.* still needs both of those
|
|
1251
|
+
earlier decisions to treat it as offloading.
|
|
1252
|
+
|
|
1253
|
+
Group offloading and on-demand residency both qualify: each leaves its component in
|
|
1254
|
+
system memory between uses, so materializing the pipeline on the device first would
|
|
1255
|
+
load in full exactly what these were configured to avoid holding.
|
|
1256
|
+
|
|
1257
|
+
Args:
|
|
1258
|
+
configuration: Configuration of the component being loaded
|
|
1259
|
+
|
|
1260
|
+
Returns:
|
|
1261
|
+
True when any per-component entry keeps its component off the device
|
|
1262
|
+
"""
|
|
1263
|
+
components = configuration.get("components") or {}
|
|
1264
|
+
return any(
|
|
1265
|
+
isinstance(settings, dict)
|
|
1266
|
+
and (
|
|
1267
|
+
settings.get("group_offload") is not None
|
|
1268
|
+
or settings.get("residency") == "on_demand"
|
|
1269
|
+
)
|
|
1270
|
+
for settings in components.values()
|
|
1271
|
+
)
|
|
1272
|
+
|
|
1273
|
+
|
|
1274
|
+
def loading_device(configuration):
|
|
1275
|
+
"""The device a component's weights are materialized on while it loads.
|
|
1276
|
+
|
|
1277
|
+
Offloading brings each part of a model onto the device only while it runs, so the
|
|
1278
|
+
weights have to land in system memory first. A default torch device pointing at the
|
|
1279
|
+
GPU would build every module directly in VRAM instead, running a large pipeline out
|
|
1280
|
+
of memory before its offload hooks are ever installed.
|
|
1281
|
+
|
|
1282
|
+
Args:
|
|
1283
|
+
configuration: Configuration of the component being loaded
|
|
1284
|
+
|
|
1285
|
+
Returns:
|
|
1286
|
+
A context manager active for the duration of the load
|
|
1287
|
+
"""
|
|
1288
|
+
offloads = (
|
|
1289
|
+
configuration.get("offload", None) is not None
|
|
1290
|
+
or configuration.get("group_offload", None) is not None
|
|
1291
|
+
or has_component_group_offload(configuration)
|
|
1292
|
+
)
|
|
1293
|
+
|
|
1294
|
+
if offloads:
|
|
1295
|
+
logger.debug("Loading into system memory - the component will be offloaded")
|
|
1296
|
+
return torch.device("cpu")
|
|
1297
|
+
|
|
1298
|
+
return contextlib.nullcontext()
|
|
1299
|
+
|
|
1300
|
+
|
|
1301
|
+
def get_block_configs(configuration, component):
|
|
1302
|
+
"""The block configs a workflow sets on a modular pipeline, checked against it.
|
|
1303
|
+
|
|
1304
|
+
A modular pipeline's blocks declare configs of their own - values they read while
|
|
1305
|
+
they run rather than components or call arguments. MiniMax-H3 declares three, and
|
|
1306
|
+
they are how the canvas the request generates on and the resolution its references
|
|
1307
|
+
are encoded at are set:
|
|
1308
|
+
|
|
1309
|
+
"configs": { "canvas_short_edge": 768, "reference_image_short_edge": 1024 }
|
|
1310
|
+
|
|
1311
|
+
This is deliberately not a knob per config per model. Every modular pipeline
|
|
1312
|
+
declares its own set, `update_components()` sets any of them, and what a workflow
|
|
1313
|
+
may say here is whatever the pipeline it named declares. The names are checked
|
|
1314
|
+
because update_components ignores the ones it does not know with a warning, and a
|
|
1315
|
+
silently dropped config reads as a setting that did nothing.
|
|
1316
|
+
|
|
1317
|
+
Args:
|
|
1318
|
+
configuration: Pipeline configuration dictionary
|
|
1319
|
+
component: The loaded pipeline the configs are for
|
|
1320
|
+
|
|
1321
|
+
Returns:
|
|
1322
|
+
Dict of config name to value, empty when the workflow sets none
|
|
1323
|
+
|
|
1324
|
+
Raises:
|
|
1325
|
+
ValueError: If the pipeline takes no configs, or does not declare one by name
|
|
1326
|
+
"""
|
|
1327
|
+
configs = configuration.get("configs", None)
|
|
1328
|
+
if not configs:
|
|
1329
|
+
return {}
|
|
1330
|
+
|
|
1331
|
+
if not has_method(component, "update_components"):
|
|
1332
|
+
raise ValueError(
|
|
1333
|
+
f"'configs' is only supported on modular pipelines, "
|
|
1334
|
+
f"{type(component).__name__} does not have update_components"
|
|
1335
|
+
)
|
|
1336
|
+
|
|
1337
|
+
# The specs a pipeline builds from its blocks. Guarded rather than indexed - a
|
|
1338
|
+
# pipeline that stops keeping them under this name should lose the check, not
|
|
1339
|
+
# the feature
|
|
1340
|
+
declared = getattr(component, "_config_specs", None)
|
|
1341
|
+
if declared is not None:
|
|
1342
|
+
unknown = [name for name in configs if name not in declared]
|
|
1343
|
+
if unknown:
|
|
1344
|
+
raise ValueError(
|
|
1345
|
+
f"{type(component).__name__} declares no config named "
|
|
1346
|
+
f"{', '.join(sorted(unknown))} - the ones it declares are "
|
|
1347
|
+
f"{', '.join(sorted(declared)) or 'none'}"
|
|
1348
|
+
)
|
|
1349
|
+
|
|
1350
|
+
logger.info(f"Setting block configs: {', '.join(configs)}")
|
|
1351
|
+
return dict(configs)
|
|
1352
|
+
|
|
1353
|
+
|
|
1354
|
+
def place_component(
|
|
1355
|
+
component, component_name, configuration, device, components_manager=None
|
|
1356
|
+
):
|
|
1357
|
+
"""Give a loaded component its offloading hooks and its device.
|
|
1358
|
+
|
|
1359
|
+
Split out of load_component because placement has to come last. Every hook
|
|
1360
|
+
here - group offloading, layerwise casting, model or sequential CPU offload -
|
|
1361
|
+
pins the modules and weights that exist when it is installed, so anything that
|
|
1362
|
+
adds or replaces weights afterwards (a LoRA, an IP-Adapter) is left outside the
|
|
1363
|
+
hook's bookkeeping: sequential offload streams the weights it recorded onto the
|
|
1364
|
+
accelerator and the adapter's own tensors are never among them, which runs the
|
|
1365
|
+
step on uninitialized weights and produces NaN.
|
|
1366
|
+
|
|
1367
|
+
Args:
|
|
1368
|
+
component: The loaded pipeline or component
|
|
1369
|
+
component_name: What is being placed, for the log
|
|
1370
|
+
configuration: The component's configuration block
|
|
1371
|
+
device: Device the component runs on
|
|
1372
|
+
components_manager: The modular pipeline's components manager, if it has one
|
|
1373
|
+
|
|
1374
|
+
Returns:
|
|
1375
|
+
The placed component
|
|
1376
|
+
"""
|
|
1377
|
+
# Handle group_offload configuration
|
|
1378
|
+
group_offload_configuration = get_group_offload_configuration(configuration, device)
|
|
1379
|
+
if group_offload_configuration is not None:
|
|
1380
|
+
component.enable_group_offload(**group_offload_configuration)
|
|
1381
|
+
|
|
1382
|
+
# Handle enable_layerwise_casting configuration
|
|
1383
|
+
enable_layerwise_casting_configuration = configuration.get(
|
|
1384
|
+
"enable_layerwise_casting", None
|
|
1385
|
+
)
|
|
1386
|
+
if enable_layerwise_casting_configuration is not None:
|
|
1387
|
+
component.enable_layerwise_casting(**enable_layerwise_casting_configuration)
|
|
1388
|
+
|
|
1389
|
+
# Configure component device settings
|
|
1390
|
+
preserve_device_placement = configuration.get("preserve_device_placement", False)
|
|
1391
|
+
offload = configuration.get("offload", None)
|
|
1392
|
+
|
|
1393
|
+
# Offloading streams a model between system memory and an accelerator - there is
|
|
1394
|
+
# nothing to stream to when the run is on the CPU
|
|
1395
|
+
if offload is not None and get_device_type(device) == "cpu":
|
|
1396
|
+
logger.warning(f"Ignoring '{offload}' offload - {device} is not an accelerator")
|
|
1397
|
+
offload = None
|
|
1398
|
+
|
|
1399
|
+
if offload == "model":
|
|
1400
|
+
logger.debug(f"Enabling model CPU offload onto {device}")
|
|
1401
|
+
component.enable_model_cpu_offload(device=device)
|
|
1402
|
+
elif offload == "sequential":
|
|
1403
|
+
logger.debug(f"Enabling sequential CPU offload onto {device}")
|
|
1404
|
+
for excluded_name in configuration.get("exclude_from_cpu_offload", []):
|
|
1405
|
+
logger.debug(f"Excluding {excluded_name} from CPU offload")
|
|
1406
|
+
component._exclude_from_cpu_offload.append(excluded_name)
|
|
1407
|
+
component.enable_sequential_cpu_offload(device=device)
|
|
1408
|
+
elif components_manager is not None and auto_cpu_offload_active(
|
|
1409
|
+
configuration, device
|
|
1410
|
+
):
|
|
1411
|
+
# Moving everything to the device here would defeat the offloading - the
|
|
1412
|
+
# manager's hooks bring each component on device as the pipeline needs it
|
|
1413
|
+
logger.debug("Device placement is owned by the components manager")
|
|
1414
|
+
elif has_component_group_offload(configuration):
|
|
1415
|
+
# configure_components() installs group-offload hooks per-component after
|
|
1416
|
+
# this returns - moving the whole pipeline to the device now would load it
|
|
1417
|
+
# in full before those hooks exist, defeating the offloading
|
|
1418
|
+
logger.info(
|
|
1419
|
+
f"components configure group offloading - not moving pipeline to {device}"
|
|
1420
|
+
)
|
|
1421
|
+
elif hasattr(component, "to") and not preserve_device_placement:
|
|
1422
|
+
logger.debug(f"Moving {component_name} to device: {device}")
|
|
1423
|
+
component = component.to(device)
|
|
1424
|
+
|
|
1425
|
+
return component
|
|
1426
|
+
|
|
1427
|
+
|
|
1428
|
+
def load_component(
|
|
1429
|
+
component_name,
|
|
1430
|
+
configuration,
|
|
1431
|
+
from_pretrained_arguments,
|
|
1432
|
+
device,
|
|
1433
|
+
reused_components=None,
|
|
1434
|
+
defer_placement=False,
|
|
1435
|
+
):
|
|
1436
|
+
"""Load and configure a pipeline or component.
|
|
1437
|
+
|
|
1438
|
+
Args:
|
|
1439
|
+
component_name: What is being loaded, for the log
|
|
1440
|
+
configuration: The component's configuration block
|
|
1441
|
+
from_pretrained_arguments: Arguments for the constructor
|
|
1442
|
+
device: Device the component is loaded for
|
|
1443
|
+
reused_components: Components an earlier step shared into this one, by name
|
|
1444
|
+
defer_placement: Load the component without placing it - the caller calls
|
|
1445
|
+
place_component once it has finished altering the weights
|
|
1446
|
+
"""
|
|
1447
|
+
component_type = configuration["component_type"]
|
|
1448
|
+
component = None
|
|
1449
|
+
|
|
1450
|
+
# A standard pipeline takes a component as a constructor argument. A modular one
|
|
1451
|
+
# cannot: it is built from the component specs in its own index and given the
|
|
1452
|
+
# objects afterwards, which is also what keeps load_components() from pulling a
|
|
1453
|
+
# second copy of the weights - it skips the components already registered
|
|
1454
|
+
reused_components = reused_components or {}
|
|
1455
|
+
takes_components_after_load = has_method(component_type, "update_components")
|
|
1456
|
+
if reused_components and not takes_components_after_load:
|
|
1457
|
+
from_pretrained_arguments.update(reused_components)
|
|
1458
|
+
|
|
1459
|
+
# A modular pipeline can hand its components to a ComponentsManager, which then
|
|
1460
|
+
# owns their device placement
|
|
1461
|
+
components_manager = create_components_manager(configuration, device)
|
|
1462
|
+
if components_manager is not None:
|
|
1463
|
+
from_pretrained_arguments["components_manager"] = components_manager
|
|
1464
|
+
|
|
1465
|
+
# MPS (Apple Silicon) has numerical instability with float16 matmul operations,
|
|
1466
|
+
# producing NaN values that result in black images. The dtype is left as asked for -
|
|
1467
|
+
# silently loading a model in a dtype the workflow did not request would be worse -
|
|
1468
|
+
# so this only warns.
|
|
1469
|
+
if (
|
|
1470
|
+
get_device_type(device) == "mps"
|
|
1471
|
+
and from_pretrained_arguments.get("torch_dtype") == torch.float16
|
|
1472
|
+
):
|
|
1473
|
+
logger.warning(
|
|
1474
|
+
f"On MPS devices float16 produces NaN values on Apple Silicon"
|
|
1475
|
+
f"Consider changing torch_dtype from float16 to float32 for {component_name} "
|
|
1476
|
+
)
|
|
1477
|
+
|
|
1478
|
+
try:
|
|
1479
|
+
with loading_device(configuration):
|
|
1480
|
+
# Load from model name
|
|
1481
|
+
if "model_name" in from_pretrained_arguments:
|
|
1482
|
+
model_name = from_pretrained_arguments.pop("model_name")
|
|
1483
|
+
logger.info(f"Loading {component_name} from model: {model_name}")
|
|
1484
|
+
component = component_type.from_pretrained(
|
|
1485
|
+
model_name, **from_pretrained_arguments
|
|
1486
|
+
)
|
|
1487
|
+
|
|
1488
|
+
# Load from single file
|
|
1489
|
+
elif "from_single_file" in from_pretrained_arguments:
|
|
1490
|
+
from_single_file = from_pretrained_arguments.pop("from_single_file")
|
|
1491
|
+
logger.info(
|
|
1492
|
+
f"Loading {component_name} from single file: {from_single_file}"
|
|
1493
|
+
)
|
|
1494
|
+
component = component_type.from_single_file(
|
|
1495
|
+
from_single_file, **from_pretrained_arguments
|
|
1496
|
+
)
|
|
1497
|
+
|
|
1498
|
+
# Create new component
|
|
1499
|
+
else:
|
|
1500
|
+
logger.info(f"Creating new {component_name}")
|
|
1501
|
+
component = component_type(**from_pretrained_arguments)
|
|
1502
|
+
|
|
1503
|
+
# Register the shared components before anything is pulled, so the
|
|
1504
|
+
# weights an earlier step already loaded and quantized are the ones
|
|
1505
|
+
# this step runs on rather than a second copy of them. The block
|
|
1506
|
+
# configs go in the same call - update_components takes both
|
|
1507
|
+
update_arguments = get_block_configs(configuration, component)
|
|
1508
|
+
if reused_components and takes_components_after_load:
|
|
1509
|
+
logger.info(
|
|
1510
|
+
f"Reusing {', '.join(reused_components)} from an earlier step"
|
|
1511
|
+
)
|
|
1512
|
+
update_arguments.update(reused_components)
|
|
1513
|
+
if update_arguments:
|
|
1514
|
+
component.update_components(**update_arguments)
|
|
1515
|
+
|
|
1516
|
+
# Modular pipelines load only their config in from_pretrained - the component
|
|
1517
|
+
# weights are pulled separately by load_components()
|
|
1518
|
+
load_components_arguments = get_load_components_arguments(configuration)
|
|
1519
|
+
if load_components_arguments is not None:
|
|
1520
|
+
if not has_method(component, "load_components"):
|
|
1521
|
+
raise ValueError(
|
|
1522
|
+
f"load_components is only supported on modular pipelines, "
|
|
1523
|
+
f"{component_type.__name__} does not have it"
|
|
1524
|
+
)
|
|
1525
|
+
logger.info(f"Loading components for {component_name}")
|
|
1526
|
+
component.load_components(**load_components_arguments)
|
|
1527
|
+
|
|
1528
|
+
if defer_placement:
|
|
1529
|
+
# The caller places this itself, once it has finished loading the
|
|
1530
|
+
# things that alter the weights - see place_component
|
|
1531
|
+
logger.debug(f"Deferring placement of {component_name}")
|
|
1532
|
+
return component
|
|
1533
|
+
|
|
1534
|
+
return place_component(
|
|
1535
|
+
component, component_name, configuration, device, components_manager
|
|
1536
|
+
)
|
|
1537
|
+
|
|
1538
|
+
except Exception as e:
|
|
1539
|
+
# One log line with the full traceback - every error class was logged
|
|
1540
|
+
# and re-raised identically
|
|
1541
|
+
logger.error(f"{type(e).__name__} loading {component_name}: {e}", exc_info=True)
|
|
1542
|
+
raise
|
|
1543
|
+
|
|
1544
|
+
|
|
1545
|
+
def apply_sdnq_optimizations(pipeline, component_names):
|
|
1546
|
+
"""Apply SDNQ quantized matmul optimization to pipeline components.
|
|
1547
|
+
|
|
1548
|
+
Uses sdnq's apply_sdnq_options_to_model to enable INT8 matmul
|
|
1549
|
+
on supported hardware (CUDA, XPU).
|
|
1550
|
+
|
|
1551
|
+
Args:
|
|
1552
|
+
pipeline: The loaded diffusers pipeline
|
|
1553
|
+
component_names: List of component names to optimize (e.g., ["transformer", "text_encoder"])
|
|
1554
|
+
"""
|
|
1555
|
+
try:
|
|
1556
|
+
from sdnq.loader import apply_sdnq_options_to_model
|
|
1557
|
+
from sdnq.common import use_torch_compile as triton_is_available
|
|
1558
|
+
except ImportError:
|
|
1559
|
+
logger.warning("sdnq not installed, skipping SDNQ optimizations")
|
|
1560
|
+
return
|
|
1561
|
+
|
|
1562
|
+
if not triton_is_available:
|
|
1563
|
+
logger.info("Triton not available, skipping SDNQ quantized matmul optimization")
|
|
1564
|
+
return
|
|
1565
|
+
|
|
1566
|
+
if not (
|
|
1567
|
+
torch.cuda.is_available() or hasattr(torch, "xpu") and torch.xpu.is_available()
|
|
1568
|
+
):
|
|
1569
|
+
logger.info(
|
|
1570
|
+
"SDNQ quantized matmul requires CUDA or XPU, skipping on this device"
|
|
1571
|
+
)
|
|
1572
|
+
return
|
|
1573
|
+
|
|
1574
|
+
for name in component_names:
|
|
1575
|
+
# A missing name (typo) and a registered-but-unloaded one both mean "nothing
|
|
1576
|
+
# to optimize here" for this call - same warn-and-skip either way
|
|
1577
|
+
try:
|
|
1578
|
+
component = get_component(pipeline, name)
|
|
1579
|
+
except ValueError:
|
|
1580
|
+
component = None
|
|
1581
|
+
|
|
1582
|
+
if component is not None:
|
|
1583
|
+
logger.info(f"Applying SDNQ quantized matmul to {name}")
|
|
1584
|
+
setattr(
|
|
1585
|
+
pipeline,
|
|
1586
|
+
name,
|
|
1587
|
+
apply_sdnq_options_to_model(component, use_quantized_matmul=True),
|
|
1588
|
+
)
|
|
1589
|
+
else:
|
|
1590
|
+
logger.warning(
|
|
1591
|
+
f"Component '{name}' not found on pipeline, skipping SDNQ optimization"
|
|
1592
|
+
)
|
|
1593
|
+
|
|
1594
|
+
|
|
1595
|
+
def get_cache_transformer(pipeline):
|
|
1596
|
+
"""Find the denoiser a cache hook attaches to.
|
|
1597
|
+
|
|
1598
|
+
Most pipelines register theirs as 'transformer', but a modular pipeline names
|
|
1599
|
+
it after the workflow it serves - MiniMax-H3's ref2va denoises through
|
|
1600
|
+
'transformer_ref'. Looking only for 'transformer' silently skips caching on
|
|
1601
|
+
those, so try the alternates diffusers' modular pipelines actually use.
|
|
1602
|
+
|
|
1603
|
+
Args:
|
|
1604
|
+
pipeline: The loaded diffusers pipeline
|
|
1605
|
+
|
|
1606
|
+
Returns:
|
|
1607
|
+
The transformer component, or None when the pipeline has none
|
|
1608
|
+
"""
|
|
1609
|
+
for name in ("transformer", "transformer_ref"):
|
|
1610
|
+
transformer = getattr(pipeline, name, None)
|
|
1611
|
+
if transformer is not None:
|
|
1612
|
+
return transformer
|
|
1613
|
+
return None
|
|
1614
|
+
|
|
1615
|
+
|
|
1616
|
+
@contextlib.contextmanager
|
|
1617
|
+
def stateful_cache_context(pipeline):
|
|
1618
|
+
"""Provide the context a stateful cache hook reads its state through.
|
|
1619
|
+
|
|
1620
|
+
first_block, mag and layer_skip keep per-context state, and their hooks go
|
|
1621
|
+
through diffusers' StateManager, which raises "No context is set" unless a
|
|
1622
|
+
context is active. A DiffusionPipeline sets one around each denoising step and
|
|
1623
|
+
clears the state afterwards in maybe_free_model_hooks; ModularPipeline is not a
|
|
1624
|
+
DiffusionPipeline and does neither, so caching a modular pipeline dies on the
|
|
1625
|
+
first step - and would otherwise carry the previous run's residuals into the
|
|
1626
|
+
next run of a pipeline this process keeps loaded.
|
|
1627
|
+
|
|
1628
|
+
One context spans the whole call rather than each step. The state is keyed by
|
|
1629
|
+
context name, so re-entering per step only re-reads the same entry. Pipelines
|
|
1630
|
+
that run separate conditional and unconditional passes name a context per pass
|
|
1631
|
+
to keep their caches apart, which a shared context would defeat - but a modular
|
|
1632
|
+
pipeline that needed that would be setting its own contexts already, and this
|
|
1633
|
+
is a no-op for pipelines whose cache is not enabled.
|
|
1634
|
+
"""
|
|
1635
|
+
transformer = get_cache_transformer(pipeline)
|
|
1636
|
+
if transformer is None or not getattr(transformer, "is_cache_enabled", False):
|
|
1637
|
+
yield
|
|
1638
|
+
return
|
|
1639
|
+
|
|
1640
|
+
logger.debug(f"Entering cache context for {transformer.__class__.__name__}")
|
|
1641
|
+
try:
|
|
1642
|
+
with transformer.cache_context(_CACHE_CONTEXT_NAME):
|
|
1643
|
+
yield
|
|
1644
|
+
finally:
|
|
1645
|
+
# Private, but it is what diffusers' own pipelines call and there is no
|
|
1646
|
+
# public equivalent. Also clears the context an errored call left set
|
|
1647
|
+
transformer._reset_stateful_cache()
|
|
1648
|
+
|
|
1649
|
+
|
|
1650
|
+
def enable_cache_on_transformer(pipeline, cache_config):
|
|
1651
|
+
"""Enable cache configuration on the pipeline's transformer.
|
|
1652
|
+
|
|
1653
|
+
Args:
|
|
1654
|
+
pipeline: The loaded diffusers pipeline
|
|
1655
|
+
cache_config: Cache configuration object from get_cache_configuration()
|
|
1656
|
+
"""
|
|
1657
|
+
transformer = get_cache_transformer(pipeline)
|
|
1658
|
+
if transformer is None:
|
|
1659
|
+
logger.warning("Pipeline has no transformer, skipping cache configuration")
|
|
1660
|
+
return
|
|
1661
|
+
|
|
1662
|
+
if not hasattr(transformer, "enable_cache"):
|
|
1663
|
+
logger.warning(
|
|
1664
|
+
f"{transformer.__class__.__name__} does not support enable_cache(), skipping"
|
|
1665
|
+
)
|
|
1666
|
+
return
|
|
1667
|
+
|
|
1668
|
+
# FasterCache decides skipping from the pipeline's current timestep. The
|
|
1669
|
+
# callback is a callable, which workflow JSON cannot express, and diffusers
|
|
1670
|
+
# calls it unconditionally on every denoiser forward - left None, the first
|
|
1671
|
+
# inference step dies. Wire it to the pipeline here, where both exist
|
|
1672
|
+
if (
|
|
1673
|
+
cache_config.__class__.__name__ == "FasterCacheConfig"
|
|
1674
|
+
and getattr(cache_config, "current_timestep_callback", None) is None
|
|
1675
|
+
):
|
|
1676
|
+
logger.debug("Wiring FasterCache current_timestep_callback to the pipeline")
|
|
1677
|
+
cache_config.current_timestep_callback = lambda: pipeline._current_timestep
|
|
1678
|
+
|
|
1679
|
+
# first_block, mag and layer_skip resolve the transformer's block class
|
|
1680
|
+
# through diffusers' registry and raise when it is absent - fill in the
|
|
1681
|
+
# blocks diffusers has not registered before handing the config over
|
|
1682
|
+
register_cache_blocks()
|
|
1683
|
+
|
|
1684
|
+
logger.info(
|
|
1685
|
+
f"Enabling {cache_config.__class__.__name__} on {transformer.__class__.__name__}"
|
|
1686
|
+
)
|
|
1687
|
+
transformer.enable_cache(cache_config)
|