diffusers-workflow 0.4.0a3__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (171) hide show
  1. diffusers_workflow-0.4.0a3.dist-info/METADATA +310 -0
  2. diffusers_workflow-0.4.0a3.dist-info/RECORD +171 -0
  3. diffusers_workflow-0.4.0a3.dist-info/WHEEL +5 -0
  4. diffusers_workflow-0.4.0a3.dist-info/entry_points.txt +6 -0
  5. diffusers_workflow-0.4.0a3.dist-info/licenses/LICENSE +201 -0
  6. diffusers_workflow-0.4.0a3.dist-info/top_level.txt +1 -0
  7. dw/__init__.py +353 -0
  8. dw/arguments.py +906 -0
  9. dw/cache_blocks.json +16 -0
  10. dw/cache_blocks.py +145 -0
  11. dw/community_pipelines/pipeline_flux_rf_inversion.py +1184 -0
  12. dw/events.py +78 -0
  13. dw/hub_cache.py +289 -0
  14. dw/introspection.py +458 -0
  15. dw/log_setup.py +45 -0
  16. dw/pipeline_processors/chain.py +750 -0
  17. dw/pipeline_processors/config_objects.py +235 -0
  18. dw/pipeline_processors/pipeline.py +1687 -0
  19. dw/pipeline_processors/remote.py +18 -0
  20. dw/previous_results.py +259 -0
  21. dw/prompt_weighting.py +378 -0
  22. dw/repl.py +298 -0
  23. dw/repl_commands.py +808 -0
  24. dw/repl_worker.py +129 -0
  25. dw/result.py +850 -0
  26. dw/run.py +92 -0
  27. dw/schema.py +24 -0
  28. dw/security.py +379 -0
  29. dw/serve.py +70 -0
  30. dw/server/__init__.py +2 -0
  31. dw/server/app.py +588 -0
  32. dw/server/jobs.py +547 -0
  33. dw/server/ui/assets/abap-08VXUWAP.js +1 -0
  34. dw/server/ui/assets/apex-BWPQTe0t.js +1 -0
  35. dw/server/ui/assets/azcli-Bc_sGQ0U.js +1 -0
  36. dw/server/ui/assets/bat-i0X4ZdIN.js +1 -0
  37. dw/server/ui/assets/bicep-B5-_aFwp.js +2 -0
  38. dw/server/ui/assets/cameligo-DMUM7wLl.js +1 -0
  39. dw/server/ui/assets/clojure-Cm7r79vr.js +1 -0
  40. dw/server/ui/assets/codicon-Brq4_Ui5.ttf +0 -0
  41. dw/server/ui/assets/coffee-Ba7i2nA0.js +1 -0
  42. dw/server/ui/assets/cpp-C7h46wYY.js +1 -0
  43. dw/server/ui/assets/csharp-BKxtCVv1.js +1 -0
  44. dw/server/ui/assets/csp-bTuwJoIa.js +1 -0
  45. dw/server/ui/assets/css-DIMkf-bt.js +3 -0
  46. dw/server/ui/assets/css.worker-B3ciXF_0.js +93 -0
  47. dw/server/ui/assets/cssMode-CEh6hWi2.js +1 -0
  48. dw/server/ui/assets/cypher-CVaqCwHa.js +1 -0
  49. dw/server/ui/assets/dart-onAF5SnQ.js +1 -0
  50. dw/server/ui/assets/dockerfile-DZFCIeNp.js +1 -0
  51. dw/server/ui/assets/ecl-D05T4iGw.js +1 -0
  52. dw/server/ui/assets/editor-jjEx9u7D.css +1 -0
  53. dw/server/ui/assets/editor.api-CExg3_mM.js +847 -0
  54. dw/server/ui/assets/editor.worker-q-txB4vs.js +30 -0
  55. dw/server/ui/assets/elixir-6RTg0lbw.js +1 -0
  56. dw/server/ui/assets/flow9-C5_-GSwl.js +1 -0
  57. dw/server/ui/assets/freemarker2-DH6orYh2.js +3 -0
  58. dw/server/ui/assets/fsharp-C8Ef5oNN.js +1 -0
  59. dw/server/ui/assets/go-C-y9NEjX.js +1 -0
  60. dw/server/ui/assets/graphql-fmXr3nnJ.js +1 -0
  61. dw/server/ui/assets/handlebars-CbrMVW4Q.js +1 -0
  62. dw/server/ui/assets/hcl-CpzslTdj.js +1 -0
  63. dw/server/ui/assets/html-YDNPZw2M.js +1 -0
  64. dw/server/ui/assets/html.worker-C93Ht9o9.js +506 -0
  65. dw/server/ui/assets/htmlMode-B_zSGWO2.js +1 -0
  66. dw/server/ui/assets/index-B7-VcYS-.css +1 -0
  67. dw/server/ui/assets/index-D_EiPU3b.js +13 -0
  68. dw/server/ui/assets/ini-sBoK_t0W.js +1 -0
  69. dw/server/ui/assets/java-BEtHBSE6.js +1 -0
  70. dw/server/ui/assets/javascript-dYuBvioq.js +1 -0
  71. dw/server/ui/assets/json.worker-B2V3pomh.js +62 -0
  72. dw/server/ui/assets/jsonMode-CUqLM39V.js +7 -0
  73. dw/server/ui/assets/julia-Bri6UV-V.js +1 -0
  74. dw/server/ui/assets/kotlin-BOotOW0E.js +1 -0
  75. dw/server/ui/assets/less-B9JPFI3C.js +2 -0
  76. dw/server/ui/assets/lexon-CfSJPG6W.js +1 -0
  77. dw/server/ui/assets/liquid-D6vxBzMv.js +1 -0
  78. dw/server/ui/assets/lspLanguageFeatures-1WJ2palX.js +4 -0
  79. dw/server/ui/assets/lua-CsQS60Ue.js +1 -0
  80. dw/server/ui/assets/m3-D-oSqn_W.js +1 -0
  81. dw/server/ui/assets/markdown-Cimd5fb3.js +1 -0
  82. dw/server/ui/assets/mdx-SHQb6vmD.js +1 -0
  83. dw/server/ui/assets/mips-CIPQ_RoX.js +1 -0
  84. dw/server/ui/assets/monaco--ixms01u.css +1 -0
  85. dw/server/ui/assets/monaco-CP-s5rcP.js +56 -0
  86. dw/server/ui/assets/msdax-DauUninz.js +1 -0
  87. dw/server/ui/assets/mysql-SOo6toE5.js +1 -0
  88. dw/server/ui/assets/objective-c-FvmIjYaQ.js +1 -0
  89. dw/server/ui/assets/pascal-DrH0SRf2.js +1 -0
  90. dw/server/ui/assets/pascaligo-D-ptJ9y-.js +1 -0
  91. dw/server/ui/assets/perl-oz_6vUea.js +1 -0
  92. dw/server/ui/assets/pgsql-DTj74zXo.js +1 -0
  93. dw/server/ui/assets/php-nr791fC2.js +1 -0
  94. dw/server/ui/assets/pla-CopQ2nXW.js +1 -0
  95. dw/server/ui/assets/postiats-43DmfD33.js +1 -0
  96. dw/server/ui/assets/powerquery-D3hlyOfw.js +1 -0
  97. dw/server/ui/assets/powershell-DmHpPYUd.js +1 -0
  98. dw/server/ui/assets/protobuf-C531GsRP.js +2 -0
  99. dw/server/ui/assets/pug-Z5eAx3Zn.js +1 -0
  100. dw/server/ui/assets/python-x0_EGHq9.js +1 -0
  101. dw/server/ui/assets/qsharp-DkqhCAOL.js +1 -0
  102. dw/server/ui/assets/r-BwWrilGY.js +1 -0
  103. dw/server/ui/assets/razor-BZC4LQDP.js +1 -0
  104. dw/server/ui/assets/redis-ClamHrr6.js +1 -0
  105. dw/server/ui/assets/redshift-DT7zqm-g.js +1 -0
  106. dw/server/ui/assets/restructuredtext-BYgofb2h.js +1 -0
  107. dw/server/ui/assets/ruby-DezsRK8O.js +1 -0
  108. dw/server/ui/assets/rust-DdL9SqIa.js +1 -0
  109. dw/server/ui/assets/sb-CcwsVR0C.js +1 -0
  110. dw/server/ui/assets/scala-DHpiXF5c.js +1 -0
  111. dw/server/ui/assets/scheme-BeGwcela.js +1 -0
  112. dw/server/ui/assets/scss-gp-XZpBa.js +3 -0
  113. dw/server/ui/assets/shell-CC2rA5mh.js +1 -0
  114. dw/server/ui/assets/solidity-BEEn4gHE.js +1 -0
  115. dw/server/ui/assets/sophia-CRfGWb83.js +1 -0
  116. dw/server/ui/assets/sparql-D_Lu-MrJ.js +1 -0
  117. dw/server/ui/assets/sql-NEE52Syq.js +1 -0
  118. dw/server/ui/assets/st-DbInun42.js +1 -0
  119. dw/server/ui/assets/swift-Bxkupp3x.js +1 -0
  120. dw/server/ui/assets/systemverilog-Bz4Y3fRF.js +1 -0
  121. dw/server/ui/assets/tcl-DISqw1ZD.js +1 -0
  122. dw/server/ui/assets/ts.worker-D7T1-Ig5.js +67738 -0
  123. dw/server/ui/assets/tsMode-BTfA6SbD.js +11 -0
  124. dw/server/ui/assets/twig-De2hgUGE.js +1 -0
  125. dw/server/ui/assets/typescript-CWA4MsNk.js +1 -0
  126. dw/server/ui/assets/typespec-B8J7ngcE.js +1 -0
  127. dw/server/ui/assets/vb-DV3o63ZY.js +1 -0
  128. dw/server/ui/assets/wgsl-DpFanUEy.js +298 -0
  129. dw/server/ui/assets/workers-CWU0uvj5.js +1 -0
  130. dw/server/ui/assets/xml-KmfTm3rg.js +1 -0
  131. dw/server/ui/assets/yaml-nFO_dDS6.js +1 -0
  132. dw/server/ui/index.html +17 -0
  133. dw/settings.py +77 -0
  134. dw/step.py +132 -0
  135. dw/tasks/audio_utils.py +266 -0
  136. dw/tasks/background_remover.py +43 -0
  137. dw/tasks/borders.py +113 -0
  138. dw/tasks/concat_videos.py +80 -0
  139. dw/tasks/depth_estimator.py +54 -0
  140. dw/tasks/diffusion_upscale.py +109 -0
  141. dw/tasks/format_messages.py +24 -0
  142. dw/tasks/gather.py +139 -0
  143. dw/tasks/image_to_text.py +43 -0
  144. dw/tasks/image_utils.py +661 -0
  145. dw/tasks/interpolate_frames.py +227 -0
  146. dw/tasks/model_cache.py +39 -0
  147. dw/tasks/pair_audio.py +58 -0
  148. dw/tasks/qr_code.py +19 -0
  149. dw/tasks/restore_faces.py +175 -0
  150. dw/tasks/rife_model.py +192 -0
  151. dw/tasks/segment.py +121 -0
  152. dw/tasks/task.py +474 -0
  153. dw/tasks/tensor_image.py +57 -0
  154. dw/tasks/text_generation.py +168 -0
  155. dw/tasks/text_sections.py +80 -0
  156. dw/tasks/upscale.py +203 -0
  157. dw/tasks/video_utils.py +154 -0
  158. dw/tasks/zoe_depth.py +71 -0
  159. dw/teacache.py +376 -0
  160. dw/teacache_models.json +99 -0
  161. dw/test.py +29 -0
  162. dw/type_helpers.py +68 -0
  163. dw/validate.py +43 -0
  164. dw/variables.py +153 -0
  165. dw/worker.py +517 -0
  166. dw/workflow.py +553 -0
  167. dw/workflow_schema.json +1157 -0
  168. dw/workflows/augment_prompt.json +65 -0
  169. dw/workflows/describe_image.json +58 -0
  170. dw/workflows/h3_context_ir.json +57 -0
  171. dw/workflows/test.json +31 -0
@@ -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)