diffusers-workflow 0.4.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (260) hide show
  1. diffusers_workflow-0.4.0.dist-info/METADATA +318 -0
  2. diffusers_workflow-0.4.0.dist-info/RECORD +260 -0
  3. diffusers_workflow-0.4.0.dist-info/WHEEL +5 -0
  4. diffusers_workflow-0.4.0.dist-info/entry_points.txt +7 -0
  5. diffusers_workflow-0.4.0.dist-info/licenses/LICENSE +201 -0
  6. diffusers_workflow-0.4.0.dist-info/top_level.txt +2 -0
  7. dw/__init__.py +440 -0
  8. dw/adapter_compatibility.py +226 -0
  9. dw/arguments.py +1231 -0
  10. dw/assessment_rules.py +159 -0
  11. dw/assets.py +130 -0
  12. dw/cache_blocks.json +16 -0
  13. dw/cache_blocks.py +146 -0
  14. dw/community_pipelines/pipeline_flux_rf_inversion.py +1184 -0
  15. dw/content_types.py +150 -0
  16. dw/dissolve_frame_errors.py +121 -0
  17. dw/docs/ACCELERATION.md +352 -0
  18. dw/docs/AGENT_LOOP.md +95 -0
  19. dw/docs/DEPENDENCIES.md +91 -0
  20. dw/docs/IP_ADAPTER.md +109 -0
  21. dw/docs/LORAS.md +131 -0
  22. dw/docs/MCP.md +517 -0
  23. dw/docs/PROMPT_WEIGHTING.md +78 -0
  24. dw/docs/QUANTIZATION.md +230 -0
  25. dw/docs/RECIPES_24GB.md +201 -0
  26. dw/docs/RELEASING.md +195 -0
  27. dw/docs/REMOTE.md +140 -0
  28. dw/docs/REPL_COMMANDS.md +121 -0
  29. dw/docs/REPL_WORKER_GUIDE.md +51 -0
  30. dw/docs/SECURITY.md +272 -0
  31. dw/docs/SECURITY_QUICKREF.md +112 -0
  32. dw/docs/SERVER.md +679 -0
  33. dw/docs/TASKS.md +1741 -0
  34. dw/docs/TESTING.md +71 -0
  35. dw/docs/WORKFLOW_GUIDE.md +2038 -0
  36. dw/docs/WORKSPACES.md +316 -0
  37. dw/download_watch.py +335 -0
  38. dw/elision.py +306 -0
  39. dw/events.py +275 -0
  40. dw/for_each.py +409 -0
  41. dw/host_memory.py +258 -0
  42. dw/host_memory_projection.py +230 -0
  43. dw/hub_cache.py +432 -0
  44. dw/introspection.py +1228 -0
  45. dw/kernel_availability.py +208 -0
  46. dw/locations.py +599 -0
  47. dw/log_setup.py +45 -0
  48. dw/loudness.py +82 -0
  49. dw/media_audio.py +217 -0
  50. dw/media_frames.py +367 -0
  51. dw/media_info.py +297 -0
  52. dw/pipeline_processors/chain.py +821 -0
  53. dw/pipeline_processors/config_objects.py +237 -0
  54. dw/pipeline_processors/pipeline.py +2297 -0
  55. dw/pipeline_processors/remote.py +46 -0
  56. dw/plan.py +920 -0
  57. dw/previous_results.py +411 -0
  58. dw/probe_paths.py +59 -0
  59. dw/prompt_schema.json +48 -0
  60. dw/prompt_weighting.py +378 -0
  61. dw/prompts.py +159 -0
  62. dw/realize.py +250 -0
  63. dw/reference_limits.py +215 -0
  64. dw/reference_names.py +125 -0
  65. dw/repl.py +338 -0
  66. dw/repl_commands.py +836 -0
  67. dw/repl_worker.py +159 -0
  68. dw/result.py +1720 -0
  69. dw/result_fps.py +82 -0
  70. dw/run.py +162 -0
  71. dw/runs.py +768 -0
  72. dw/scalar_result_validation.py +97 -0
  73. dw/schema.py +283 -0
  74. dw/security.py +1038 -0
  75. dw/select_validation.py +115 -0
  76. dw/serve.py +277 -0
  77. dw/server/__init__.py +2 -0
  78. dw/server/app.py +4586 -0
  79. dw/server/assess.py +132 -0
  80. dw/server/catalog_shape.py +487 -0
  81. dw/server/enhancers.py +129 -0
  82. dw/server/exports.py +480 -0
  83. dw/server/guides.py +257 -0
  84. dw/server/jobs.py +1561 -0
  85. dw/server/mcp_mount.py +95 -0
  86. dw/server/netinfo.py +124 -0
  87. dw/server/observed_cost.py +379 -0
  88. dw/server/sysinfo.py +71 -0
  89. dw/server/ui/assets/abap-08VXUWAP.js +1 -0
  90. dw/server/ui/assets/apex-BWPQTe0t.js +1 -0
  91. dw/server/ui/assets/azcli-Bc_sGQ0U.js +1 -0
  92. dw/server/ui/assets/bat-i0X4ZdIN.js +1 -0
  93. dw/server/ui/assets/bicep-B5-_aFwp.js +2 -0
  94. dw/server/ui/assets/cameligo-DMUM7wLl.js +1 -0
  95. dw/server/ui/assets/clojure-Cm7r79vr.js +1 -0
  96. dw/server/ui/assets/codicon-Brq4_Ui5.ttf +0 -0
  97. dw/server/ui/assets/coffee-Ba7i2nA0.js +1 -0
  98. dw/server/ui/assets/cpp-C7h46wYY.js +1 -0
  99. dw/server/ui/assets/csharp-BKxtCVv1.js +1 -0
  100. dw/server/ui/assets/csp-bTuwJoIa.js +1 -0
  101. dw/server/ui/assets/css-DIMkf-bt.js +3 -0
  102. dw/server/ui/assets/css.worker-B3ciXF_0.js +93 -0
  103. dw/server/ui/assets/cssMode-CPznxfY8.js +1 -0
  104. dw/server/ui/assets/cypher-CVaqCwHa.js +1 -0
  105. dw/server/ui/assets/dart-onAF5SnQ.js +1 -0
  106. dw/server/ui/assets/dockerfile-DZFCIeNp.js +1 -0
  107. dw/server/ui/assets/ecl-D05T4iGw.js +1 -0
  108. dw/server/ui/assets/editor-jjEx9u7D.css +1 -0
  109. dw/server/ui/assets/editor.api-CpWcotrd.js +847 -0
  110. dw/server/ui/assets/editor.worker-q-txB4vs.js +30 -0
  111. dw/server/ui/assets/elixir-6RTg0lbw.js +1 -0
  112. dw/server/ui/assets/flow9-C5_-GSwl.js +1 -0
  113. dw/server/ui/assets/freemarker2-CXtRM8N4.js +3 -0
  114. dw/server/ui/assets/fsharp-C8Ef5oNN.js +1 -0
  115. dw/server/ui/assets/go-C-y9NEjX.js +1 -0
  116. dw/server/ui/assets/graphql-fmXr3nnJ.js +1 -0
  117. dw/server/ui/assets/handlebars-N7x-6NMY.js +1 -0
  118. dw/server/ui/assets/hcl-CpzslTdj.js +1 -0
  119. dw/server/ui/assets/html-PhsdjHSr.js +1 -0
  120. dw/server/ui/assets/html.worker-C93Ht9o9.js +506 -0
  121. dw/server/ui/assets/htmlMode-Dgj0SEok.js +1 -0
  122. dw/server/ui/assets/index-3Vw6WAPW.css +1 -0
  123. dw/server/ui/assets/index-DgrYhQd9.js +43 -0
  124. dw/server/ui/assets/ini-sBoK_t0W.js +1 -0
  125. dw/server/ui/assets/java-BEtHBSE6.js +1 -0
  126. dw/server/ui/assets/javascript-BJqN9Qhv.js +1 -0
  127. dw/server/ui/assets/json.worker-B2V3pomh.js +62 -0
  128. dw/server/ui/assets/jsonMode-DbM4SWSv.js +7 -0
  129. dw/server/ui/assets/julia-Bri6UV-V.js +1 -0
  130. dw/server/ui/assets/kotlin-BOotOW0E.js +1 -0
  131. dw/server/ui/assets/less-B9JPFI3C.js +2 -0
  132. dw/server/ui/assets/lexon-CfSJPG6W.js +1 -0
  133. dw/server/ui/assets/liquid-BWr8lEc4.js +1 -0
  134. dw/server/ui/assets/lspLanguageFeatures-C1iGuDyZ.js +4 -0
  135. dw/server/ui/assets/lua-CsQS60Ue.js +1 -0
  136. dw/server/ui/assets/m3-D-oSqn_W.js +1 -0
  137. dw/server/ui/assets/markdown-Cimd5fb3.js +1 -0
  138. dw/server/ui/assets/mdx-DAdMi_0p.js +1 -0
  139. dw/server/ui/assets/mips-CIPQ_RoX.js +1 -0
  140. dw/server/ui/assets/monaco--ixms01u.css +1 -0
  141. dw/server/ui/assets/monaco-BGCeEqaw.js +56 -0
  142. dw/server/ui/assets/msdax-DauUninz.js +1 -0
  143. dw/server/ui/assets/mysql-SOo6toE5.js +1 -0
  144. dw/server/ui/assets/objective-c-FvmIjYaQ.js +1 -0
  145. dw/server/ui/assets/pascal-DrH0SRf2.js +1 -0
  146. dw/server/ui/assets/pascaligo-D-ptJ9y-.js +1 -0
  147. dw/server/ui/assets/perl-oz_6vUea.js +1 -0
  148. dw/server/ui/assets/pgsql-DTj74zXo.js +1 -0
  149. dw/server/ui/assets/php-nr791fC2.js +1 -0
  150. dw/server/ui/assets/pla-CopQ2nXW.js +1 -0
  151. dw/server/ui/assets/postiats-43DmfD33.js +1 -0
  152. dw/server/ui/assets/powerquery-D3hlyOfw.js +1 -0
  153. dw/server/ui/assets/powershell-DmHpPYUd.js +1 -0
  154. dw/server/ui/assets/protobuf-C531GsRP.js +2 -0
  155. dw/server/ui/assets/pug-Z5eAx3Zn.js +1 -0
  156. dw/server/ui/assets/python-Bcn70HdC.js +1 -0
  157. dw/server/ui/assets/qsharp-DkqhCAOL.js +1 -0
  158. dw/server/ui/assets/r-BwWrilGY.js +1 -0
  159. dw/server/ui/assets/razor-D1HmNnby.js +1 -0
  160. dw/server/ui/assets/redis-ClamHrr6.js +1 -0
  161. dw/server/ui/assets/redshift-DT7zqm-g.js +1 -0
  162. dw/server/ui/assets/restructuredtext-BYgofb2h.js +1 -0
  163. dw/server/ui/assets/ruby-DezsRK8O.js +1 -0
  164. dw/server/ui/assets/rust-DdL9SqIa.js +1 -0
  165. dw/server/ui/assets/sb-CcwsVR0C.js +1 -0
  166. dw/server/ui/assets/scala-DHpiXF5c.js +1 -0
  167. dw/server/ui/assets/scheme-BeGwcela.js +1 -0
  168. dw/server/ui/assets/scss-gp-XZpBa.js +3 -0
  169. dw/server/ui/assets/shell-CC2rA5mh.js +1 -0
  170. dw/server/ui/assets/solidity-BEEn4gHE.js +1 -0
  171. dw/server/ui/assets/sophia-CRfGWb83.js +1 -0
  172. dw/server/ui/assets/sparql-D_Lu-MrJ.js +1 -0
  173. dw/server/ui/assets/sql-NEE52Syq.js +1 -0
  174. dw/server/ui/assets/st-DbInun42.js +1 -0
  175. dw/server/ui/assets/swift-Bxkupp3x.js +1 -0
  176. dw/server/ui/assets/systemverilog-Bz4Y3fRF.js +1 -0
  177. dw/server/ui/assets/tcl-DISqw1ZD.js +1 -0
  178. dw/server/ui/assets/ts.worker-D7T1-Ig5.js +67738 -0
  179. dw/server/ui/assets/tsMode-D6u0XmOW.js +11 -0
  180. dw/server/ui/assets/twig-De2hgUGE.js +1 -0
  181. dw/server/ui/assets/typescript-BU6v-LMV.js +1 -0
  182. dw/server/ui/assets/typespec-B8J7ngcE.js +1 -0
  183. dw/server/ui/assets/vb-DV3o63ZY.js +1 -0
  184. dw/server/ui/assets/wgsl-DpFanUEy.js +298 -0
  185. dw/server/ui/assets/workers-Cn7cTUKr.js +1 -0
  186. dw/server/ui/assets/xml--0LP2Lwk.js +1 -0
  187. dw/server/ui/assets/yaml-mpBg9jnt.js +1 -0
  188. dw/server/ui/index.html +17 -0
  189. dw/server/updater.py +192 -0
  190. dw/settings.py +98 -0
  191. dw/shot_span_preflight.py +116 -0
  192. dw/shots.py +359 -0
  193. dw/slice_preflight.py +148 -0
  194. dw/step.py +187 -0
  195. dw/step_cache.py +442 -0
  196. dw/subfolders.py +107 -0
  197. dw/task_domains.py +307 -0
  198. dw/tasks/assess.py +826 -0
  199. dw/tasks/audio_transcription.py +88 -0
  200. dw/tasks/audio_utils.py +1862 -0
  201. dw/tasks/background_remover.py +43 -0
  202. dw/tasks/borders.py +113 -0
  203. dw/tasks/compose_text.py +74 -0
  204. dw/tasks/concat_videos.py +300 -0
  205. dw/tasks/depth_estimator.py +54 -0
  206. dw/tasks/diffusion_upscale.py +109 -0
  207. dw/tasks/dissolve_videos.py +342 -0
  208. dw/tasks/format_messages.py +24 -0
  209. dw/tasks/gather.py +173 -0
  210. dw/tasks/grade.py +97 -0
  211. dw/tasks/image_to_text.py +43 -0
  212. dw/tasks/image_utils.py +764 -0
  213. dw/tasks/interpolate_frames.py +252 -0
  214. dw/tasks/judge.py +68 -0
  215. dw/tasks/model_cache.py +55 -0
  216. dw/tasks/pair_audio.py +268 -0
  217. dw/tasks/qr_code.py +19 -0
  218. dw/tasks/restore_faces.py +175 -0
  219. dw/tasks/rife_model.py +192 -0
  220. dw/tasks/segment.py +121 -0
  221. dw/tasks/select.py +111 -0
  222. dw/tasks/speech_generation.py +228 -0
  223. dw/tasks/stabilize.py +129 -0
  224. dw/tasks/task.py +920 -0
  225. dw/tasks/tensor_image.py +57 -0
  226. dw/tasks/text_generation.py +169 -0
  227. dw/tasks/text_sections.py +80 -0
  228. dw/tasks/upscale.py +203 -0
  229. dw/tasks/video_utils.py +624 -0
  230. dw/tasks/zoe_depth.py +71 -0
  231. dw/teacache.py +381 -0
  232. dw/teacache_models.json +99 -0
  233. dw/test.py +29 -0
  234. dw/type_helpers.py +231 -0
  235. dw/validate.py +68 -0
  236. dw/variable_constraints.py +444 -0
  237. dw/variables.py +443 -0
  238. dw/video_extensions.py +141 -0
  239. dw/vram_estimate.py +116 -0
  240. dw/worker.py +764 -0
  241. dw/workflow.py +2007 -0
  242. dw/workflow_schema.json +1346 -0
  243. dw/workflow_sources.py +383 -0
  244. dw/workflows/h3_context_ir.json +57 -0
  245. dw/workflows/test.json +31 -0
  246. dw/workspace.py +730 -0
  247. dw_mcp/__init__.py +6 -0
  248. dw_mcp/__main__.py +133 -0
  249. dw_mcp/assets.py +336 -0
  250. dw_mcp/authoring.py +114 -0
  251. dw_mcp/catalog.py +360 -0
  252. dw_mcp/client.py +486 -0
  253. dw_mcp/diagnose.py +371 -0
  254. dw_mcp/exports.py +84 -0
  255. dw_mcp/guides.py +35 -0
  256. dw_mcp/media.py +638 -0
  257. dw_mcp/models.py +97 -0
  258. dw_mcp/prompts.py +104 -0
  259. dw_mcp/server.py +1343 -0
  260. dw_mcp/workspaces.py +212 -0
@@ -0,0 +1,2297 @@
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 ..security import (
20
+ require_trusted_from_pretrained_arguments,
21
+ require_trusted_pre_load_modules,
22
+ )
23
+ from .. import empty_device_cache, get_device_type, resolve_device
24
+ from diffusers import attention_backend
25
+
26
+ # dw.prompt_weighting (transformers) and diffusers.hooks (peft, bitsandbytes) are
27
+ # imported where they are used - at module scope they add seconds to every startup
28
+
29
+ from ..events import WorkflowCancelled, emit_phase, emit_warning, get_context
30
+ from .. import download_watch
31
+ from huggingface_hub.errors import HfHubHTTPError
32
+
33
+ logger = logging.getLogger("dw")
34
+
35
+ # Names the cache state a run accumulates. Any stable string works - it only has to
36
+ # match itself across the steps of one call
37
+ _CACHE_CONTEXT_NAME = "dw"
38
+
39
+ optional_component_names = [
40
+ "controlnet",
41
+ "transformer",
42
+ "transformer_2",
43
+ "vae",
44
+ "unet",
45
+ "text_encoder",
46
+ "text_encoder_2",
47
+ "text_encoder_3",
48
+ "tokenizer",
49
+ "tokenizer_2",
50
+ "tokenizer_3",
51
+ "image_encoder",
52
+ "feature_extractor",
53
+ "prompt_enhancer_head",
54
+ "model",
55
+ ]
56
+
57
+ # Pipeline-definition keys that can never name a component
58
+ _NON_COMPONENT_KEYS = {
59
+ "configuration",
60
+ "from_pretrained_arguments",
61
+ "arguments",
62
+ "scheduler",
63
+ "loras",
64
+ "ip_adapter",
65
+ "seed",
66
+ "remote_text_encoder",
67
+ }
68
+
69
+
70
+ def declared_component_names(pipeline_definition):
71
+ """The component names a pipeline definition can load or configure.
72
+
73
+ The known names plus any other key shaped like a component - a dict carrying
74
+ 'from_pretrained_arguments' (a scheduler carries 'from_config_args' instead).
75
+ Diffusers grows new component names faster than the list above; a workflow
76
+ naming one gets it loaded rather than silently dropped.
77
+ """
78
+ names = list(optional_component_names)
79
+ for key, value in pipeline_definition.items():
80
+ if (
81
+ key not in names
82
+ and key not in _NON_COMPONENT_KEYS
83
+ and isinstance(value, dict)
84
+ and "from_pretrained_arguments" in value
85
+ ):
86
+ logger.info(f"Treating '{key}' as a component definition")
87
+ names.append(key)
88
+ return names
89
+
90
+
91
+ class Pipeline:
92
+ """
93
+ Manages pipeline initialization, configuration, and execution.
94
+ Handles loading of models, schedulers, and adapters.
95
+ """
96
+
97
+ def __init__(
98
+ self,
99
+ pipeline_definition,
100
+ default_seed,
101
+ device,
102
+ pipeline=None,
103
+ output_dir=None,
104
+ file_prefix=None,
105
+ ):
106
+ """
107
+ Initialize pipeline with configuration and device settings.
108
+
109
+ Args:
110
+ pipeline_definition: Dictionary containing pipeline configuration
111
+ default_seed: Seed value for reproducibility
112
+ device: Device to run pipeline on (e.g., 'cuda', 'mps', 'cpu') - the
113
+ configuration's own 'device' takes precedence over it
114
+ pipeline: Optional existing pipeline to use
115
+ output_dir: The workflow's output directory - where a chained run
116
+ with save_segments writes its segment files
117
+ file_prefix: Naming prefix for those files, matching the step's
118
+ result naming (workflow id + step name)
119
+ """
120
+ self.pipeline_definition = pipeline_definition
121
+ self.default_seed = default_seed
122
+ # A step can pin itself to a device, overriding the one dw is running on. It
123
+ # becomes the default for this pipeline's components as well, and is
124
+ # translated to a backend this machine has so the workflow travels
125
+ self.device = resolve_device(self.configuration.get("device", device))
126
+ self.pipeline = pipeline
127
+ self.output_dir = output_dir
128
+ self.file_prefix = file_prefix
129
+ # What a chained run calls the segment it is on, so progress can say
130
+ # which one the denoise counter belongs to - it restarts per segment
131
+ self.segment_label = None
132
+ logger.debug(f"Initialized pipeline with device: {self.device}")
133
+
134
+ @property
135
+ def configuration(self):
136
+ return self.pipeline_definition.get("configuration", {})
137
+
138
+ @property
139
+ def name(self):
140
+ return self.from_pretrained_arguments.get("model_name", "")
141
+
142
+ @property
143
+ def from_pretrained_arguments(self):
144
+ return self.pipeline_definition.get("from_pretrained_arguments", {})
145
+
146
+ @property
147
+ def argument_template(self):
148
+ return self.pipeline_definition["arguments"]
149
+
150
+ def component_names(self, key):
151
+ """The component names one of the sharing lists holds.
152
+
153
+ The lists were only ever read off the pipeline itself, while the schema and
154
+ the guide put them in its configuration - a workflow written to the docs
155
+ shared nothing and said nothing about it. Both places are read now.
156
+
157
+ Args:
158
+ key: 'shared_components' or 'reused_components'
159
+
160
+ Returns:
161
+ List of component names
162
+ """
163
+ return list(self.pipeline_definition.get(key, [])) + list(
164
+ self.configuration.get(key, [])
165
+ )
166
+
167
+ def resolve_reused_components(self, shared_components):
168
+ """The components an earlier step shared that this one asks to reuse.
169
+
170
+ Args:
171
+ shared_components: Dictionary of components shared between pipelines
172
+
173
+ Returns:
174
+ Dict of component name to the component itself
175
+
176
+ Raises:
177
+ ValueError: If a name was never shared by an earlier step
178
+ """
179
+ reused = {}
180
+ for name in self.component_names("reused_components"):
181
+ if name not in shared_components:
182
+ raise ValueError(
183
+ f"Cannot reuse component '{name}' - no earlier step shared it. "
184
+ f"Shared so far: {sorted(shared_components) or 'nothing'}"
185
+ )
186
+ logger.debug(f"Reusing component: {name}")
187
+ reused[name] = shared_components[name]
188
+ return reused
189
+
190
+ def populate_from_pretrained_arguments(self, device, shared_components):
191
+ """
192
+ Prepare arguments for pipeline creation, including shared components.
193
+
194
+ The loaded components go into a copy, not into the definition they were read
195
+ from. The definition belongs to the workflow and outlives every step, so a
196
+ component stored there is a component the run holds until it ends: releasing
197
+ the pipeline frees nothing, and a workflow that loads a second large model
198
+ after releasing the first runs out of memory holding both. Copying also
199
+ leaves the definition intact for a second load - load_component consumes
200
+ 'model_name' out of the arguments it is handed.
201
+
202
+ Args:
203
+ device: Device to run pipeline on
204
+ shared_components: Dictionary of components shared between pipelines
205
+ """
206
+ logger.debug("Populating from_pretrained arguments")
207
+ from_pretrained_arguments = dict(self.from_pretrained_arguments)
208
+
209
+ # Load optional components (controlnet, vae, unet, etc.), including any
210
+ # component-shaped key outside the known names
211
+ for component_name in declared_component_names(self.pipeline_definition):
212
+ self.load_optional_component(
213
+ component_name, from_pretrained_arguments, device
214
+ )
215
+
216
+ # Handle remote text encoder configuration by setting local text_encoder to None
217
+ if self.pipeline_definition.get("remote_text_encoder", None):
218
+ logger.info("Configuring remote text encoder")
219
+ from_pretrained_arguments["text_encoder"] = None
220
+
221
+ return from_pretrained_arguments
222
+
223
+ def check_trusted(self):
224
+ """Refuse an untrusted definition before anything says a load began.
225
+
226
+ The gates themselves live inside load() and load_component(), which is
227
+ where they have to be - that is the last point before the bytes are
228
+ fetched. But `load()` is entered under a 'loading' phase event, so a
229
+ run refused by them emitted the same marker as one that loaded a model
230
+ and then failed, and job events could no longer tell the two apart
231
+ (#137). This runs the same checks over the definition first, so the
232
+ caller emits 'loading' only once a load can actually begin. It is a
233
+ pre-flight, not the boundary: the in-load checks stay.
234
+
235
+ Raises:
236
+ UntrustedWorkflowError: If untrusted and the definition reaches
237
+ for remote code
238
+ """
239
+ require_trusted_pre_load_modules(self.configuration.get("pre_load_modules", []))
240
+ # Walk the whole definition rather than the component names this class
241
+ # knows: a gate that only covers what it remembers to enumerate stops
242
+ # covering a block added later
243
+ self._check_trusted_block(self.pipeline_definition, "pipeline")
244
+
245
+ @staticmethod
246
+ def _check_trusted_block(block, what):
247
+ if not isinstance(block, dict):
248
+ return
249
+ require_trusted_from_pretrained_arguments(
250
+ block.get("from_pretrained_arguments"), what
251
+ )
252
+ for key, value in block.items():
253
+ if key == "from_pretrained_arguments":
254
+ continue
255
+ if isinstance(value, dict):
256
+ Pipeline._check_trusted_block(value, key)
257
+ elif isinstance(value, list):
258
+ for entry in value:
259
+ Pipeline._check_trusted_block(entry, key)
260
+
261
+ def load(self, shared_components):
262
+ """
263
+ Load and configure the pipeline with all components.
264
+
265
+ Args:
266
+ shared_components: Dictionary of components shared between pipelines
267
+ """
268
+ logger.debug(f"Loading pipeline: {self.name}")
269
+
270
+ # Import modules that need to register with diffusers/transformers before loading
271
+ # (e.g., sdnq registers its quantization method on import). This runs
272
+ # arbitrary python at import time, so an untrusted workflow is refused
273
+ # here unless --trust-workflows was passed - see docs/SECURITY.md
274
+ pre_load_modules = self.configuration.get("pre_load_modules", [])
275
+ require_trusted_pre_load_modules(pre_load_modules)
276
+ for module_name in pre_load_modules:
277
+ logger.info(f"Pre-loading module: {module_name}")
278
+ importlib.import_module(module_name)
279
+
280
+ # A load that raises partway has already built some of what it was
281
+ # asked for - the pipeline itself, its quantized weights, a placement
282
+ # half applied - and none of that is reachable from the caller, which
283
+ # never received a pipeline. Left alone it stays resident until the
284
+ # exception is handled and something else happens to collect, so a
285
+ # retry loads its own copy on top of the last attempt's: three failures
286
+ # is three pipelines' worth of dead weight. Tear the attempt down here
287
+ # instead, then let the failure carry on
288
+ try:
289
+ # Prepare arguments and load pipeline
290
+ from_pretrained_arguments = self.populate_from_pretrained_arguments(
291
+ self.device, shared_components
292
+ )
293
+ reused_components = self.resolve_reused_components(shared_components)
294
+
295
+ # Adapters add weights to the components they attach to, and an offloading
296
+ # hook only streams the weights that existed when it was installed - so a
297
+ # pipeline that loads any is placed after they are on it, not at load
298
+ adapters_to_load = bool(self.pipeline_definition.get("loras", [])) or (
299
+ self.pipeline_definition.get("ip_adapter", None) is not None
300
+ )
301
+
302
+ # Load and configure the main pipeline
303
+ self.pipeline = load_component(
304
+ "pipeline",
305
+ self.configuration,
306
+ from_pretrained_arguments,
307
+ self.device,
308
+ reused_components,
309
+ defer_placement=adapters_to_load,
310
+ )
311
+
312
+ # Enable attention slicing if explicitly requested or automatically on MPS
313
+ # MPS benefits from slicing since Metal shares system RAM with the GPU
314
+ if self.configuration.get("enable_attention_slicing", False) or (
315
+ get_device_type(self.device) == "mps"
316
+ and not self.configuration.get("disable_attention_slicing", False)
317
+ ):
318
+ # Modular pipelines have no attention slicing - on MPS this is applied
319
+ # automatically, so skip rather than fail when the pipeline lacks it
320
+ if has_method(self.pipeline, "enable_attention_slicing"):
321
+ logger.debug("Enabling attention slicing for pipeline")
322
+ self.pipeline.enable_attention_slicing()
323
+ else:
324
+ logger.debug(
325
+ f"{type(self.pipeline).__name__} does not support attention slicing, skipping"
326
+ )
327
+
328
+ # configure components that are not shared
329
+ self.configure_loaded_components()
330
+
331
+ # Apply SDNQ quantized matmul optimization to specified components
332
+ sdnq_optimize = self.configuration.get("sdnq_optimize", [])
333
+ if sdnq_optimize:
334
+ apply_sdnq_optimizations(self.pipeline, sdnq_optimize)
335
+
336
+ # Enable diffusers built-in cache acceleration on transformer
337
+ cache_config = get_cache_configuration(self.configuration)
338
+ if cache_config is not None:
339
+ enable_cache_on_transformer(self.pipeline, cache_config)
340
+
341
+ # Configure the schedulers if specified - a pipeline that denoises two
342
+ # modalities against two schedules configures each of them separately
343
+ load_and_configure_scheduler(
344
+ self.pipeline_definition.get("scheduler", None), self.pipeline
345
+ )
346
+ load_and_configure_scheduler(
347
+ self.pipeline_definition.get("audio_scheduler", None),
348
+ self.pipeline,
349
+ "audio_scheduler",
350
+ )
351
+
352
+ self.publish_shared_components(shared_components)
353
+
354
+ # Load and configure LoRA models
355
+ load_loras(self.pipeline_definition.get("loras", []), self.pipeline)
356
+
357
+ # Load and configure IP-Adapter
358
+ load_ip_adapter(
359
+ self.pipeline_definition.get("ip_adapter", None), self.pipeline
360
+ )
361
+
362
+ # The adapters are on the pipeline now, so its offloading hooks can be
363
+ # installed over the weights they added
364
+ if adapters_to_load:
365
+ self.pipeline = place_component(
366
+ self.pipeline,
367
+ "pipeline",
368
+ self.configuration,
369
+ self.device,
370
+ # A modular pipeline's manager owns its placement, and the manager
371
+ # load_component gave it is the one it holds
372
+ getattr(self.pipeline, "_components_manager", None),
373
+ )
374
+
375
+ # Place the components the pipeline loaded itself, once everything that alters
376
+ # them - dtypes, adapters, quantized matmuls - has been applied. Offloading hooks
377
+ # installed before those would be fighting them
378
+ configure_components(
379
+ self.pipeline, self.configuration, self.device, reused_components
380
+ )
381
+
382
+ # Set up random generator if needed - no_generator is a boolean, so an
383
+ # explicit false still gets a generator
384
+ if not self.configuration.get("no_generator", False):
385
+ logger.debug("Setting up random generator")
386
+ self.argument_template["generator"] = torch.Generator(
387
+ self.device
388
+ ).manual_seed(self.pipeline_definition.get("seed", self.default_seed))
389
+
390
+ # Hand the first run a clean allocator. Loading churns the device even
391
+ # when little of the pipeline stays there - a quantization pass with
392
+ # 'quantization_device' set works on the accelerator and returns the
393
+ # weights to the host, and group offloading moves components off it
394
+ # again - and the cached blocks left behind are the wrong shape for
395
+ # inference. workflow.py does this between steps; a one-step workflow
396
+ # would otherwise run its only step on top of the loading debris
397
+ gc.collect()
398
+ empty_device_cache()
399
+ except BaseException:
400
+ self._discard_failed_load(shared_components)
401
+ raise
402
+
403
+ logger.debug("Pipeline loaded successfully")
404
+
405
+ def _discard_failed_load(self, shared_components):
406
+ """Drop everything a load that raised had built, and reclaim it.
407
+
408
+ Unpublishes as well as releases: a component shared before the failure
409
+ would otherwise be handed to a later step as a component of a pipeline
410
+ that does not exist. A name this pipeline reused rather than loaded
411
+ stays published - that entry belongs to the earlier pipeline that put
412
+ it there, which is still alive.
413
+ """
414
+ logger.info(f"Load of pipeline '{self.name}' failed - releasing what it built")
415
+ reused = set(self.component_names("reused_components"))
416
+ for shared_component_name in self.component_names("shared_components"):
417
+ if shared_component_name not in reused:
418
+ shared_components.pop(shared_component_name, None)
419
+ self.pipeline = None
420
+ self.argument_template.pop("generator", None)
421
+ gc.collect()
422
+ empty_device_cache()
423
+
424
+ def publish_shared_components(self, shared_components):
425
+ """Store components that will be shared with other pipelines.
426
+
427
+ Called from load, and again by the workflow when a cached pipeline is
428
+ reused - a cache hit skips load entirely, and the shared_components
429
+ dict is fresh every run, so a warm sharing step must republish or a
430
+ later reusing step finds nothing. get_component rather than getattr -
431
+ a modular pipeline registers a component it did not load as None, and
432
+ sharing that None silently would surface as a missing-component error
433
+ inside the step that reused it.
434
+ """
435
+ for shared_component_name in self.component_names("shared_components"):
436
+ component = get_component(self.pipeline, shared_component_name)
437
+ if component is None:
438
+ raise ValueError(
439
+ f"Cannot share component '{shared_component_name}' - "
440
+ f"{type(self.pipeline).__name__} registers it but has not "
441
+ f"loaded it"
442
+ )
443
+ logger.debug(f"Storing shared component: {shared_component_name}")
444
+ shared_components[shared_component_name] = component
445
+
446
+ @torch.inference_mode()
447
+ def run(self, arguments, previous_pipelines={}):
448
+ """
449
+ Execute the pipeline with given arguments.
450
+
451
+ Args:
452
+ arguments: Dictionary of arguments for pipeline execution
453
+ previous_pipelines: Dictionary of previously created pipelines
454
+
455
+ Returns:
456
+ Pipeline output or dictionary containing special outputs
457
+ """
458
+ if self.pipeline is None:
459
+ logger.error("Pipeline not initialized")
460
+ raise ValueError(
461
+ "Pipeline has not been initialized. Call load(device_identifier, shared_components) first."
462
+ )
463
+
464
+ logger.debug(f"Running pipeline with arguments: {arguments}")
465
+
466
+ try:
467
+ # Handle inversion pipeline
468
+ if self.configuration.get("inversion", False):
469
+ logger.debug("Running inversion pipeline")
470
+ invert_arguments = copy.deepcopy(arguments)
471
+ invert_arguments.pop("generator", None)
472
+ inverted_latents, image_latents, latent_image_ids = (
473
+ self.pipeline.invert(**invert_arguments)
474
+ )
475
+ return {
476
+ "inverted_latents": inverted_latents,
477
+ "image_latents": image_latents,
478
+ "latent_image_ids": latent_image_ids,
479
+ }
480
+
481
+ # Handle generation pipeline
482
+ if self.configuration.get("generate", False):
483
+ logger.debug("Running generation pipeline")
484
+ return {"generated_ids": self.pipeline.generate(**arguments)}
485
+
486
+ chain_definition = self.pipeline_definition.get("chain", None)
487
+ if chain_definition is not None:
488
+ from .chain import run_chain
489
+
490
+ logger.debug("Running chained pipeline")
491
+ return run_chain(self, chain_definition, arguments)
492
+
493
+ return self._run_once(arguments)
494
+
495
+ except WorkflowCancelled:
496
+ logger.info("Pipeline run cancelled")
497
+ raise
498
+ except Exception as e:
499
+ # One log line with the full traceback - every error class was
500
+ # logged and re-raised identically
501
+ logger.error(f"{type(e).__name__} running pipeline: {e}", exc_info=True)
502
+ raise
503
+
504
+ def _run_once(self, arguments):
505
+ """Run one standard pipeline invocation with fully resolved arguments.
506
+
507
+ This is the whole per-call execution path - prompt encoding, the
508
+ pipeline call itself, and output normalization - shared by the single
509
+ run and every segment of a chained run.
510
+ """
511
+ if self.pipeline_definition.get("remote_text_encoder", None) is not None:
512
+ logger.info("Invoking remote text encoder")
513
+ remote_config = self.pipeline_definition["remote_text_encoder"]
514
+ prompt_embeds = remote_text_encoder(
515
+ arguments.pop("prompt"),
516
+ remote_config.get("url"),
517
+ device=self.device,
518
+ )
519
+ arguments["prompt_embeds"] = prompt_embeds
520
+ elif self.configuration.get("prompt_weighting", False):
521
+ from ..prompt_weighting import apply_prompt_weighting
522
+
523
+ # The step's device override travels with the call - embeddings
524
+ # must land where the transformer runs
525
+ apply_prompt_weighting(self.pipeline, arguments, self.device)
526
+
527
+ # Run standard pipeline
528
+ logger.debug("Running standard pipeline")
529
+ output = self._execute_pipeline(arguments)
530
+
531
+ # A raw tensor result - latents, embeddings - is held for the rest of the
532
+ # workflow, so it rests in system memory instead of occupying the
533
+ # accelerator that the next step needs. Pipelines consuming it place it back
534
+ # on their own device.
535
+ if hasattr(output, "to"):
536
+ logger.debug("Moving tensor output to system memory")
537
+ output = output.to("cpu")
538
+
539
+ attach_audio_sample_rate(self.pipeline, output)
540
+ warn_if_safety_checker_blanked(output)
541
+
542
+ return output
543
+
544
+ def _execute_pipeline(self, arguments):
545
+ """Execute the pipeline with optional TeaCache and attention backend contexts."""
546
+ teacache_config = self.configuration.get("teacache", None)
547
+ attn_backend = self.configuration.get("attention_backend", None)
548
+
549
+ # Determine the execution context
550
+ if teacache_config is not None:
551
+ num_steps = arguments.get("num_inference_steps", None)
552
+ if num_steps is None:
553
+ logger.warning(
554
+ "TeaCache requires num_inference_steps in arguments, running without TeaCache"
555
+ )
556
+ return self._call_pipeline(arguments, attn_backend)
557
+
558
+ rel_l1_thresh = teacache_config.get("rel_l1_thresh", None)
559
+ coefficients = teacache_config.get("coefficients", None)
560
+ variant = teacache_config.get("variant", None)
561
+ with teacache_context(
562
+ self.pipeline, num_steps, rel_l1_thresh, coefficients, variant
563
+ ):
564
+ return self._call_pipeline(arguments, attn_backend)
565
+ else:
566
+ return self._call_pipeline(arguments, attn_backend)
567
+
568
+ def _call_pipeline(self, arguments, attn_backend):
569
+ """Call the pipeline with optional attention backend and cache contexts."""
570
+ arguments = self._with_step_callback(arguments)
571
+ # The load is over and the denoise loop is starting. Pipelines whose
572
+ # signature has no step callback report nothing else at all, so this
573
+ # is the only thing that distinguishes running from still loading
574
+ emit_phase("generating", detail=self.segment_label or self.name)
575
+ with contextlib.ExitStack() as stack:
576
+ if attn_backend is not None:
577
+ logger.info(f"Using attention backend: {attn_backend}")
578
+ stack.enter_context(attention_backend(attn_backend))
579
+
580
+ stack.enter_context(stateful_cache_context(self.pipeline))
581
+ if not self._takes_step_callback():
582
+ # A modular pipeline takes no step callback at all, so this
583
+ # is the only per-step signal it has: the denoise blocks
584
+ # drive a tqdm bar, and a bar that reports each advance is
585
+ # the difference between a slow run and a hung one
586
+ stack.enter_context(reported_progress_bars(self.pipeline))
587
+ # The bar only covers the denoise loop; the blocks around it
588
+ # are where a reference encode's minutes go (#95)
589
+ stack.enter_context(
590
+ reported_blocks(self.pipeline, self.segment_label or self.name)
591
+ )
592
+
593
+ return self.pipeline(**arguments)
594
+
595
+ def _takes_step_callback(self):
596
+ """Whether this pipeline names `callback_on_step_end` in its own
597
+ signature. Only a pipeline that names the parameter explicitly gets
598
+ one - a **kwargs signature is no promise the pipeline honors it, and
599
+ a ModularPipeline (H3, LTX-2, Qwen-Image) has no such parameter at
600
+ all, which is why it needs the progress-bar route instead."""
601
+ try:
602
+ parameters = inspect.signature(self.pipeline.__call__).parameters
603
+ except (TypeError, ValueError):
604
+ return False
605
+ return "callback_on_step_end" in parameters
606
+
607
+ def _with_step_callback(self, arguments):
608
+ """Inject a callback_on_step_end that reports per-step progress to the
609
+ active run context and raises when the run has been cancelled.
610
+
611
+ Workflow JSON cannot express a callable, so this is the only way a
612
+ diffusion-step callback ever reaches a pipeline call.
613
+ """
614
+ if not self._takes_step_callback():
615
+ return arguments
616
+
617
+ run_context = get_context()
618
+ num_steps = arguments.get("num_inference_steps", None)
619
+
620
+ def on_step_end(pipe, step_index, timestep, callback_kwargs):
621
+ total = getattr(pipe, "_num_timesteps", None) or num_steps
622
+ run_context.emit("pipeline_step", step=step_index + 1, total_steps=total)
623
+ # Past the last step the pipeline still has to decode the latents,
624
+ # which on video is minutes with the bar sitting at 100%
625
+ if total is not None and step_index + 1 >= total:
626
+ emit_phase("decoding")
627
+ run_context.check_cancelled()
628
+ return callback_kwargs
629
+
630
+ return {**arguments, "callback_on_step_end": on_step_end}
631
+
632
+ def load_optional_component(
633
+ self, component_name, from_pretrained_arguments, default_device
634
+ ):
635
+ """Load an optional component if specified in pipeline definition."""
636
+ component_definition = self.pipeline_definition.get(component_name, None)
637
+
638
+ if component_definition is not None:
639
+ logger.info(f"Loading component: {component_name}")
640
+ component_configuration = component_definition.get("configuration", None)
641
+ if component_configuration is not None:
642
+ # A copy for the same reason the pipeline's own arguments are copied:
643
+ # what goes in here is consumed by load_component, and the definition
644
+ # is the workflow's, not this load's
645
+ component_from_pretrained_arguments = dict(
646
+ component_definition["from_pretrained_arguments"]
647
+ )
648
+
649
+ # Handle quantization configuration
650
+ quantization_configuration = get_quantization_configuration(
651
+ component_definition
652
+ )
653
+ if quantization_configuration is not None:
654
+ logger.debug(f"Adding quantization config for {component_name}")
655
+ component_from_pretrained_arguments["quantization_config"] = (
656
+ quantization_configuration
657
+ )
658
+
659
+ device = resolve_device(
660
+ component_configuration.get("device", default_device)
661
+ )
662
+
663
+ component = load_component(
664
+ component_name,
665
+ component_configuration,
666
+ component_from_pretrained_arguments,
667
+ device,
668
+ )
669
+
670
+ logger.debug(f"Loaded optional component: {component_name}")
671
+ from_pretrained_arguments[component_name] = component
672
+
673
+ def configure_loaded_components(self):
674
+ # Configure VAE settings
675
+ vae = self.configuration.get("vae", {})
676
+ if vae.get("enable_slicing", False):
677
+ logger.debug("Enabling VAE slicing")
678
+ self.pipeline.vae.enable_slicing()
679
+ if vae.get("enable_tiling", False):
680
+ logger.debug("Enabling VAE tiling")
681
+ self.pipeline.vae.enable_tiling()
682
+ if vae.get("channels_last", False):
683
+ logger.debug("Setting VAE memory format")
684
+ self.pipeline.vae.to(memory_format=torch.channels_last)
685
+
686
+ # Configure UNet settings
687
+ unet = self.configuration.get("unet", {})
688
+ if unet.get("enable_forward_chunking", False):
689
+ logger.debug("Enabling UNet forward chunking")
690
+ self.pipeline.unet.enable_forward_chunking()
691
+ if unet.get("channels_last", False):
692
+ logger.debug("Setting UNet memory format")
693
+ self.pipeline.unet.to(memory_format=torch.channels_last)
694
+
695
+ # Configure UNet attention processor
696
+ if unet.get("attn_processor_type", None) is not None:
697
+ logger.debug("Enabling UNet custom attention processor")
698
+ attn_processor = unet["attn_processor_type"]()
699
+ self.pipeline.unet.set_attn_processor(attn_processor)
700
+
701
+ # Configure transformer settings
702
+ transformer = self.configuration.get("transformer", {})
703
+ if transformer.get("attn_processor_type", None) is not None:
704
+ logger.debug("Enabling transformer custom attention processor")
705
+ attn_processor = transformer["attn_processor_type"]()
706
+ self.pipeline.transformer.set_attn_processor(attn_processor)
707
+
708
+ # configure optional components
709
+ for component_name in declared_component_names(self.pipeline_definition):
710
+ component_configuration = self.configuration.get(component_name, None)
711
+ if component_configuration is None:
712
+ continue
713
+
714
+ # get_component() raises on a genuinely missing attribute (a typo) and
715
+ # returns None for one that is registered but unloaded - both cases are
716
+ # unconfigurable, so both are skipped here exactly as a plain missing
717
+ # component always was
718
+ try:
719
+ component = get_component(self.pipeline, component_name)
720
+ except ValueError:
721
+ component = None
722
+
723
+ if component is not None:
724
+ logger.debug(f"Configuring optional component: {component_name}")
725
+ torch_dtype = component_configuration.get("torch_dtype", None)
726
+ if torch_dtype is not None:
727
+ logger.debug(f"Setting {component_name} torch dtype: {torch_dtype}")
728
+ component.to(torch_dtype)
729
+
730
+
731
+ def configure_components(pipeline, configuration, default_device, reused_components=()):
732
+ """Place the components a pipeline loaded for itself.
733
+
734
+ A modular pipeline pulls its own component weights, so they are only reachable once
735
+ the pipeline is loaded - too late for the offloading load_component sets up. Group
736
+ offloading a component here streams it between system memory and the accelerator a
737
+ piece at a time, which is what fits a pipeline whose components are each larger than
738
+ the device.
739
+
740
+ A component this step reused is not one it loaded: it already carries the placement
741
+ the step that shared it gave it, and offloading hooks do not survive being applied
742
+ twice. Those are skipped, so a workflow can reuse a component into a step whose
743
+ configuration was written for loading it.
744
+
745
+ Args:
746
+ pipeline: The loaded pipeline
747
+ configuration: Pipeline configuration dictionary
748
+ default_device: Device the pipeline runs on
749
+ reused_components: Names of the components an earlier step shared into this one
750
+ """
751
+ for component_name, component_configuration in configuration.get(
752
+ "components", {}
753
+ ).items():
754
+ # A dotted path reaches inside a component, and it is the component itself
755
+ # that was shared - 'text_encoder.model' belongs to a reused 'text_encoder'
756
+ if component_name.split(".")[0] in reused_components:
757
+ logger.info(
758
+ f"Component '{component_name}' was shared by an earlier step - "
759
+ "keeping the placement that step gave it"
760
+ )
761
+ continue
762
+
763
+ component = get_component(pipeline, component_name)
764
+ if component is None:
765
+ # Registered but unloaded (e.g. a components map reused across workflow
766
+ # selections, or a component diffusers warned-and-skipped past at load) -
767
+ # skip just this entry rather than aborting the whole run
768
+ logger.warning(
769
+ f"Component '{component_name}' is not loaded (workflow selection "
770
+ "may not use it) - skipping its configuration"
771
+ )
772
+ continue
773
+
774
+ # Prune before any hooks are installed - an offload hook pins and tracks
775
+ # exactly the modules that exist when it is applied, so pruning afterwards
776
+ # would leave it streaming weights that can never run
777
+ truncate_module_lists(component, component_name, component_configuration)
778
+ replace_modules_with_identity(
779
+ component, component_name, component_configuration
780
+ )
781
+
782
+ group_offload_configuration = get_group_offload_configuration(
783
+ component_configuration, default_device
784
+ )
785
+ if group_offload_configuration is not None:
786
+ # apply_group_offloading rather than the component's own
787
+ # enable_group_offload - a component may be a transformers model, or a
788
+ # module inside one, and only diffusers models have the method
789
+ from diffusers.hooks import apply_group_offloading
790
+
791
+ logger.info(f"Group offloading {component_name}")
792
+ apply_group_offloading(component, **group_offload_configuration)
793
+
794
+ # Tiled decoding, for a component that decodes but is not the one called
795
+ # 'vae' - LTX-2.5's diffusion decoder, which decodes the whole video volume
796
+ # in one allocation unless it is told to tile
797
+ enable_tiling(component, component_name, component_configuration)
798
+
799
+ # The attention processor this component runs, for a component that is
800
+ # neither the unet nor the transformer (both covered by their own
801
+ # pipeline-level blocks) - LTX-2.5's diffusion decoder, whose default
802
+ # processor is a portable fallback rather than the path it was trained to run
803
+ set_attn_processor(component, component_name, component_configuration)
804
+
805
+ device = resolve_device(component_configuration.get("device", None))
806
+ residency = component_configuration.get("residency", "resident")
807
+ if residency == "on_demand":
808
+ apply_on_demand_placement(
809
+ component,
810
+ component_name,
811
+ device if device is not None else default_device,
812
+ group_offload_configuration is not None,
813
+ )
814
+ elif device is not None:
815
+ logger.info(f"Moving {component_name} to device: {device}")
816
+ component.to(device)
817
+
818
+ # A compiled component should pin its attention backend - the per-call
819
+ # attention_backend context manager would switch implementations under a
820
+ # compiled graph and force a recompile on every run
821
+ component_attention_backend = component_configuration.get(
822
+ "attention_backend", None
823
+ )
824
+ if component_attention_backend is not None:
825
+ logger.info(
826
+ f"Setting {component_name} attention backend: {component_attention_backend}"
827
+ )
828
+ component.set_attention_backend(component_attention_backend)
829
+
830
+ # Compile last - the graph must capture final dtypes, adapters,
831
+ # quantization, and offload hooks
832
+ compile_configuration = component_configuration.get("compile", None)
833
+ if compile_configuration is not None:
834
+ apply_compile(
835
+ component,
836
+ component_name,
837
+ compile_configuration,
838
+ device if device is not None else default_device,
839
+ )
840
+
841
+
842
+ def enable_tiling(component, component_name, component_configuration):
843
+ """Turn on tiled decoding for a component whose configuration asks for it.
844
+
845
+ The pipeline-level `vae` block covers the component actually named 'vae'. This
846
+ covers any other component that decodes - LTX-2.5's `diffusion_decoder`, which
847
+ otherwise decodes the whole video volume in a single allocation and asks for
848
+ tens of GiB at 2x resolutions. `true` takes the model's own default tile size;
849
+ a dict passes the tile and stride sizes through, which is what a card smaller
850
+ than those defaults needs.
851
+
852
+ Args:
853
+ component: The loaded component
854
+ component_name: Its name, for logging and errors
855
+ component_configuration: That component's configuration block
856
+
857
+ Raises:
858
+ ValueError: If the component has no enable_tiling() to call
859
+ """
860
+ tiling = component_configuration.get("enable_tiling", False)
861
+ if not tiling:
862
+ return
863
+
864
+ if not has_method(component, "enable_tiling"):
865
+ raise ValueError(
866
+ f"'{component_name}' does not support tiling - "
867
+ f"{type(component).__name__} has no enable_tiling()"
868
+ )
869
+
870
+ arguments = tiling if isinstance(tiling, dict) else {}
871
+ logger.info(
872
+ f"Enabling tiling on {component_name}"
873
+ + (f" with {', '.join(arguments)}" if arguments else "")
874
+ )
875
+ component.enable_tiling(**arguments)
876
+
877
+
878
+ def set_attn_processor(component, component_name, component_configuration):
879
+ """Swap a component's attention processor for the one its configuration names.
880
+
881
+ The pipeline-level `unet` and `transformer` blocks cover those two components.
882
+ This covers any other one that carries attention - LTX-2.5's
883
+ `diffusion_decoder`, whose default `LTX2VideoVaeNeighborhoodAttnProcessor` is a
884
+ portable FlexAttention fallback rather than the NATTEN path the decoder was
885
+ built around, and which diffusers' own docstring calls larger than device memory
886
+ at production grids.
887
+
888
+ The value is a type, resolved by the `_type` suffix convention, and is
889
+ constructed with no arguments - the same shape the `unet`/`transformer` blocks
890
+ have used since they were written.
891
+
892
+ Args:
893
+ component: The loaded component
894
+ component_name: Its name, for logging and errors
895
+ component_configuration: That component's configuration block
896
+
897
+ Raises:
898
+ ValueError: If the component has no set_attn_processor() to call
899
+ """
900
+ attn_processor_type = component_configuration.get("attn_processor_type", None)
901
+ if attn_processor_type is None:
902
+ return
903
+
904
+ if not has_method(component, "set_attn_processor"):
905
+ raise ValueError(
906
+ f"'{component_name}' does not take an attention processor - "
907
+ f"{type(component).__name__} has no set_attn_processor()"
908
+ )
909
+
910
+ logger.info(
911
+ f"Setting {component_name} attention processor: {attn_processor_type.__name__}"
912
+ )
913
+ component.set_attn_processor(attn_processor_type())
914
+
915
+
916
+ def _resolve_submodule(component, component_name, path):
917
+ """Follow a dotted path from a component to a module inside it.
918
+
919
+ Args:
920
+ component: The component the path starts from
921
+ component_name: Its name, for errors
922
+ path: Dotted attribute path relative to the component, e.g.
923
+ 'language_model.layers'
924
+
925
+ Returns:
926
+ (parent, attribute_name, module) - the module and where it hangs
927
+
928
+ Raises:
929
+ ValueError: If any step of the path is not an attribute
930
+ """
931
+ parent = None
932
+ module = component
933
+ for attribute_name in path.split("."):
934
+ parent = module
935
+ module = getattr(parent, attribute_name, _MISSING)
936
+ if module is _MISSING:
937
+ raise ValueError(
938
+ f"'{component_name}' has no module at '{path}' - "
939
+ f"{type(parent).__name__} has no attribute '{attribute_name}'"
940
+ )
941
+ return parent, path.rsplit(".", 1)[-1], module
942
+
943
+
944
+ def truncate_module_lists(component, component_name, component_configuration):
945
+ """Drop the tail of a ModuleList a run never reads.
946
+
947
+ An encoder used for its hidden states can run layers whose output nothing
948
+ consumes: MiniMax-H3 conditions on hidden_states[50] of its 64-layer
949
+ Qwen3-VL, so layers 51-63 compute - and, offloaded, stream from system
950
+ memory - for nothing, on every encode. Keeping 51 layers leaves
951
+ hidden_states[50] bit-identical (index 50 of the returned tuple is the
952
+ input to layer 50, recorded before it runs; keeping only 50 would make it
953
+ the final-norm output instead, which is a different tensor).
954
+
955
+ The configuration maps a dotted path inside the component to the number of
956
+ entries to keep:
957
+
958
+ "truncate_layers": { "language_model.layers": 51 }
959
+
960
+ Truncation is in place, so the component's registration on its pipeline and
961
+ its config are untouched - a block that validates against
962
+ config.num_hidden_layers still sees the checkpoint's own count.
963
+
964
+ Args:
965
+ component: The loaded component
966
+ component_name: Its name, for logging and errors
967
+ component_configuration: That component's configuration block
968
+
969
+ Raises:
970
+ ValueError: If a path does not lead to a ModuleList, or keep is not
971
+ a positive count
972
+ """
973
+ truncations = component_configuration.get("truncate_layers", None)
974
+ if not truncations:
975
+ return
976
+
977
+ for path, keep in truncations.items():
978
+ _, _, module_list = _resolve_submodule(component, component_name, path)
979
+ if not isinstance(module_list, torch.nn.ModuleList):
980
+ raise ValueError(
981
+ f"'{component_name}' cannot truncate '{path}' - it is a "
982
+ f"{type(module_list).__name__}, not a ModuleList"
983
+ )
984
+ keep = int(keep)
985
+ if keep < 1:
986
+ raise ValueError(
987
+ f"'{component_name}' truncate_layers keeps {keep} of '{path}' - "
988
+ "at least one layer has to remain"
989
+ )
990
+ if keep >= len(module_list):
991
+ logger.warning(
992
+ f"'{component_name}' truncate_layers keeps {keep} of '{path}', "
993
+ f"which already has {len(module_list)} - nothing to drop"
994
+ )
995
+ continue
996
+ logger.info(
997
+ f"Truncating {component_name} '{path}' from {len(module_list)} "
998
+ f"layers to {keep}"
999
+ )
1000
+ del module_list[keep:]
1001
+
1002
+
1003
+ def replace_modules_with_identity(component, component_name, component_configuration):
1004
+ """Swap out modules a run never calls, freeing what they hold.
1005
+
1006
+ For weights that exist on the checkpoint but are outside the path the
1007
+ pipeline actually runs - a language-model head on a model used as an
1008
+ encoder, say. The module is replaced with an Identity so the model keeps
1009
+ its shape for anything that looks the attribute up, while its parameters
1010
+ are dropped rather than held (and, offloaded, pinned) for a call that
1011
+ never comes.
1012
+
1013
+ "remove_modules": [ "lm_head" ]
1014
+
1015
+ Args:
1016
+ component: The loaded component
1017
+ component_name: Its name, for logging and errors
1018
+ component_configuration: That component's configuration block
1019
+
1020
+ Raises:
1021
+ ValueError: If a named module does not exist on the component
1022
+ """
1023
+ for path in component_configuration.get("remove_modules", []):
1024
+ parent, attribute_name, _ = _resolve_submodule(component, component_name, path)
1025
+ logger.info(f"Replacing {component_name} '{path}' with Identity")
1026
+ setattr(parent, attribute_name, torch.nn.Identity())
1027
+
1028
+
1029
+ # The calls that mean "this component is working now". A component is moved to
1030
+ # the accelerator around whichever of these it actually defines
1031
+ _ON_DEMAND_ENTRY_POINTS = ("forward", "encode", "decode")
1032
+
1033
+
1034
+ def apply_on_demand_placement(
1035
+ component, component_name, device, group_offloaded, offload_device="cpu"
1036
+ ):
1037
+ """Keep a component in system memory and move it to the device only while it runs.
1038
+
1039
+ Sits between the two placements dw already has. A 'device' component is resident
1040
+ for the whole run, which wastes the accelerator on something used twice; group
1041
+ offloading streams per submodule forward, which restreams the whole model once
1042
+ per call of every leaf - ruinous for a VAE, whose tiled decode calls its blocks
1043
+ once per tile. This moves the model as a whole around each entry point, so a
1044
+ tiling loop sits inside a single pair of transfers.
1045
+
1046
+ That trade only pays for components called a handful of times per run. A
1047
+ denoising transformer is called once per step, so per-call transfers would cost
1048
+ far more than they save - group offloading is the tool for those.
1049
+
1050
+ Args:
1051
+ component: The component to place
1052
+ component_name: Name of the component, for logging
1053
+ device: Device to run the component on
1054
+ group_offloaded: Whether group offloading was applied to this component
1055
+ offload_device: Where the component rests between calls
1056
+
1057
+ Raises:
1058
+ ValueError: If the component is also group offloaded
1059
+ """
1060
+ if group_offloaded:
1061
+ raise ValueError(
1062
+ f"Component '{component_name}' sets both 'group_offload' and "
1063
+ "'residency: on_demand'. A group offloaded module holds one group at a "
1064
+ "time and ignores the whole-model moves on-demand placement makes, so "
1065
+ "the two cannot both own its placement - pick one"
1066
+ )
1067
+
1068
+ if get_device_type(device) == "cpu":
1069
+ # Nothing to move it off of, so the wrappers would be pure overhead
1070
+ logger.debug(
1071
+ f"Ignoring 'residency: on_demand' for {component_name} - {device} is the "
1072
+ "device it would rest on anyway"
1073
+ )
1074
+ return
1075
+
1076
+ component.to(offload_device)
1077
+
1078
+ # One depth counter for the whole component, not one per entry point: decode()
1079
+ # calls forward() internally, and an inner return must not offload the model
1080
+ # out from under the call that is still running
1081
+ state = {"depth": 0}
1082
+
1083
+ def wrap(entry_point):
1084
+ original = getattr(component, entry_point, None)
1085
+ if not callable(original):
1086
+ return False
1087
+
1088
+ @functools.wraps(original)
1089
+ def on_demand(*args, **kwargs):
1090
+ if state["depth"] == 0:
1091
+ component.to(device)
1092
+ state["depth"] += 1
1093
+ try:
1094
+ return original(*args, **kwargs)
1095
+ finally:
1096
+ state["depth"] -= 1
1097
+ if state["depth"] == 0:
1098
+ component.to(offload_device)
1099
+ # Hand the freed space back to the driver rather than leaving it
1100
+ # reserved - the headroom is the entire point of doing this
1101
+ empty_device_cache()
1102
+
1103
+ # functools.wraps carries __wrapped__, so inspect.signature() still reports
1104
+ # the real parameters. Callers introspect them: MiniMax H3's denoiser picks
1105
+ # which arguments to pass by reading signature(transformer.forward)
1106
+ setattr(component, entry_point, on_demand)
1107
+ return True
1108
+
1109
+ wrapped = [name for name in _ON_DEMAND_ENTRY_POINTS if wrap(name)]
1110
+ if not wrapped:
1111
+ raise ValueError(
1112
+ f"Component '{component_name}' sets 'residency: on_demand' but defines "
1113
+ f"none of {', '.join(_ON_DEMAND_ENTRY_POINTS)}, so there is no call to "
1114
+ "move it around"
1115
+ )
1116
+ logger.info(
1117
+ f"Placing {component_name} on demand: resting on {offload_device}, "
1118
+ f"running on {device} around {', '.join(wrapped)}"
1119
+ )
1120
+
1121
+
1122
+ def apply_compile(component, component_name, compile_configuration, device):
1123
+ """Compile a component with torch.compile.
1124
+
1125
+ Compilation happens in place (nn.Module.compile) so the module stays registered
1126
+ on its pipeline. With 'repeated_blocks' true, only the model's repeated block
1127
+ classes are compiled (diffusers' regional compilation) - near the same speedup
1128
+ as full compilation with a fraction of the cold-start cost.
1129
+
1130
+ Args:
1131
+ component: The component to compile
1132
+ component_name: Name of the component, for logging
1133
+ compile_configuration: Dict of options - 'repeated_blocks' selects regional
1134
+ compilation, everything else ('mode', 'fullgraph', 'dynamic', ...) is
1135
+ passed to torch.compile
1136
+ device: Device the component runs on
1137
+ """
1138
+ # Inductor support on MPS is too immature to be worth the compile time
1139
+ if get_device_type(device) == "mps":
1140
+ logger.warning(
1141
+ f"torch.compile is not supported on MPS, skipping {component_name}"
1142
+ )
1143
+ return
1144
+
1145
+ options = dict(compile_configuration)
1146
+ repeated_blocks = options.pop("repeated_blocks", False)
1147
+
1148
+ if repeated_blocks:
1149
+ if not has_method(component, "compile_repeated_blocks"):
1150
+ raise ValueError(
1151
+ f"repeated_blocks compilation requires a diffusers model with "
1152
+ f"repeated block support, {type(component).__name__} does not have it"
1153
+ )
1154
+ logger.info(f"Compiling repeated blocks of {component_name}")
1155
+ component.compile_repeated_blocks(**options)
1156
+ else:
1157
+ logger.info(f"Compiling {component_name}")
1158
+ component.compile(**options)
1159
+
1160
+
1161
+ _MISSING = object()
1162
+
1163
+
1164
+ def get_component(pipeline, component_name):
1165
+ """Look a component up on a pipeline, by name or by a dotted path into it.
1166
+
1167
+ A dotted path reaches a module inside a component, which is how a component that
1168
+ holds the model rather than being one - a transformers model wrapping its own - is
1169
+ offloaded.
1170
+
1171
+ A modular pipeline registers a component it has not loaded (a workflow selection
1172
+ that does not use it, or one diffusers warned-and-skipped past) as a None-valued
1173
+ attribute rather than omitting it entirely - that is a real attribute, not a typo,
1174
+ so it is returned as None rather than raising. A missing attribute is still a hard
1175
+ error: it means the name itself is wrong. Callers decide what "unloaded" should
1176
+ mean for them (skip with a warning, skip silently, ...); this just tells them apart.
1177
+
1178
+ Args:
1179
+ pipeline: The loaded pipeline
1180
+ component_name: Name of the component, e.g. 'vae' or 'text_encoder.model'
1181
+
1182
+ Returns:
1183
+ The named component, or None if it (or a step along a dotted path) is
1184
+ registered but not loaded
1185
+
1186
+ Raises:
1187
+ ValueError: If the pipeline has no attribute by that name (or dotted path)
1188
+ """
1189
+ component = pipeline
1190
+ for attribute_name in component_name.split("."):
1191
+ component = getattr(component, attribute_name, _MISSING)
1192
+ if component is _MISSING:
1193
+ raise ValueError(
1194
+ f"{type(pipeline).__name__} has no component '{component_name}'"
1195
+ )
1196
+ if component is None:
1197
+ return None
1198
+
1199
+ return component
1200
+
1201
+
1202
+ def warn_if_safety_checker_blanked(output):
1203
+ """Say so when a safety checker replaced a generated image with a black one.
1204
+
1205
+ Stable Diffusion 1.5's checker false-positives readily, and it returns a
1206
+ solid black image rather than an error. Run to run that reads as the seed
1207
+ having no effect - the same result every time - so the reason belongs
1208
+ where the identical images do: the run's warnings, which is the only
1209
+ place a consumer of a `succeeded` job would ever see it.
1210
+
1211
+ Args:
1212
+ output: The pipeline output, which may carry nsfw_content_detected
1213
+ """
1214
+ flags = getattr(output, "nsfw_content_detected", None)
1215
+ if not flags:
1216
+ return
1217
+
1218
+ blanked = sum(1 for flag in flags if flag)
1219
+ if blanked:
1220
+ # emit_warning rather than logger.warning: a blanked image is a
1221
+ # succeeded job whose file is solid black, and a consumer over the
1222
+ # API or MCP sees the job's warnings list and nothing else - the log
1223
+ # line never reaches the one party that cannot tell the picture apart
1224
+ # from a rendered one (#133)
1225
+ emit_warning(
1226
+ f"The safety checker blanked {blanked} of {len(flags)} generated "
1227
+ "images - they are solid black, and no seed will change that. "
1228
+ "Pass 'safety_checker': null in from_pretrained_arguments to "
1229
+ "load the pipeline without it.",
1230
+ kind="safety_checker_blanked",
1231
+ blanked=blanked,
1232
+ images=len(flags),
1233
+ )
1234
+
1235
+
1236
+ # Where an audio pipeline's components record the rate they generate at, in
1237
+ # the order they are tried. LTX-2's vocoder names it output_sampling_rate,
1238
+ # AudioLDM2's vocoder and StableAudio's VAE name it sampling_rate
1239
+ _SAMPLE_RATE_SOURCES = (
1240
+ ("vocoder", "output_sampling_rate"),
1241
+ ("vocoder", "sampling_rate"),
1242
+ ("vae", "sampling_rate"),
1243
+ )
1244
+
1245
+
1246
+ def _component_sample_rate(pipeline, has_audios):
1247
+ for component_name, attribute in _SAMPLE_RATE_SOURCES:
1248
+ if component_name == "vae" and not has_audios:
1249
+ # A video VAE's config is not where an audio rate lives - only an
1250
+ # audio-only pipeline's VAE (StableAudio) reports one there
1251
+ continue
1252
+ config = getattr(getattr(pipeline, component_name, None), "config", None)
1253
+ sample_rate = getattr(config, attribute, None)
1254
+ if sample_rate is not None:
1255
+ return sample_rate
1256
+ return None
1257
+
1258
+
1259
+ def attach_audio_sample_rate(pipeline, output):
1260
+ """Record the generating component's sample rate on an output that carries audio.
1261
+
1262
+ Pipelines that generate audio - with their video (LTX-2, on `.audio`) or by
1263
+ itself (AudioLDM2, StableAudio, on `.audios`) - return the waveform without
1264
+ its sample rate; only the vocoder or VAE that produced it knows that. Saving
1265
+ needs the rate, so it travels with the output, and get_artifact_list wraps
1266
+ `.audios` items in an AudioTrack that carries it.
1267
+
1268
+ Args:
1269
+ pipeline: The pipeline that produced the output
1270
+ output: The pipeline output
1271
+ """
1272
+ has_audio = getattr(output, "audio", None) is not None
1273
+ has_audios = getattr(output, "audios", None) is not None
1274
+ if not (has_audio or has_audios):
1275
+ return
1276
+
1277
+ sample_rate = _component_sample_rate(pipeline, has_audios)
1278
+ if sample_rate is None:
1279
+ logger.warning(
1280
+ "Pipeline generated audio but no component reports its sample rate - "
1281
+ "set 'sample_rate' (audio results) or 'audio_sample_rate' (muxed video) "
1282
+ "in the step result to save it correctly"
1283
+ )
1284
+ return
1285
+
1286
+ logger.debug(f"Generated audio has a sample rate of {sample_rate}Hz")
1287
+ output.audio_sample_rate = sample_rate
1288
+
1289
+
1290
+ def set_adapter_alpha(pipeline, adapter_name, alpha):
1291
+ """Override the network alpha an adapter was loaded with.
1292
+
1293
+ peft scales an adapter by `scale * alpha / rank`, and the alpha comes from
1294
+ the checkpoint: a per-module `.alpha` tensor, else a `__metadata__` alpha
1295
+ where the loader honors one, else the rank itself. That is the right
1296
+ default, and for some files it is wrong. The 768p MiniMax-H3 turbo LoRAs
1297
+ record `alpha: 8` in their `__metadata__` at rank 128, which diffusers
1298
+ honors, while upstream's own 768p invocation passes `--lora-alpha 128` -
1299
+ sixteen times the strength the file asks for. The number that makes a
1300
+ distilled checkpoint hit its trained schedule is a property of the model,
1301
+ so the workflow states it rather than the engine guessing.
1302
+
1303
+ Set before the caller's set_adapters(), which is what recomputes each
1304
+ layer's scaling from the alpha found here.
1305
+
1306
+ Args:
1307
+ pipeline: The loaded pipeline
1308
+ adapter_name: Which adapter's alpha to override
1309
+ alpha: The network alpha to use
1310
+
1311
+ Raises:
1312
+ ValueError: if alpha is not positive, or no loaded layer carries the
1313
+ adapter - an alpha that silently applied to nothing is the quality
1314
+ failure it exists to prevent
1315
+ """
1316
+ alpha = float(alpha)
1317
+ if alpha <= 0:
1318
+ raise ValueError(f"A lora 'alpha' must be positive, got {alpha}.")
1319
+
1320
+ touched = 0
1321
+ ranks = set()
1322
+ for module in _lora_layers(pipeline):
1323
+ if adapter_name in getattr(module, "lora_alpha", {}):
1324
+ module.lora_alpha[adapter_name] = alpha
1325
+ ranks.add(module.r[adapter_name])
1326
+ touched += 1
1327
+
1328
+ if not touched:
1329
+ raise ValueError(
1330
+ f"Cannot set alpha {alpha} on adapter '{adapter_name}' - no loaded "
1331
+ "layer carries it. Nothing would have been scaled."
1332
+ )
1333
+
1334
+ # The figure upstream's own runner prints, because alpha alone says
1335
+ # nothing without the rank it is divided by
1336
+ effective = sorted(alpha / rank for rank in ranks)
1337
+ logger.info(
1338
+ f"Adapter '{adapter_name}': alpha {alpha} over rank(s) "
1339
+ f"{sorted(ranks)} - scaling {effective[0]:.6g}"
1340
+ + (f" to {effective[-1]:.6g}" if len(effective) > 1 else "")
1341
+ + " before the adapter weight"
1342
+ )
1343
+
1344
+
1345
+ def _lora_layers(pipeline):
1346
+ """Every peft-wrapped module under a pipeline, whatever holds it.
1347
+
1348
+ A pipeline names the components a LoRA can reach; a modular one holds them
1349
+ in `components` instead. Walking both and de-duplicating by identity is
1350
+ what keeps this working for MiniMax-H3, whose two DiT partitions are
1351
+ separate components, without enumerating model names here.
1352
+ """
1353
+ from peft.tuners.tuners_utils import BaseTunerLayer
1354
+
1355
+ seen = set()
1356
+ holders = []
1357
+ for name in getattr(pipeline, "_lora_loadable_modules", []) or []:
1358
+ holders.append(getattr(pipeline, name, None))
1359
+ if not any(holder is not None for holder in holders):
1360
+ components = getattr(pipeline, "components", None)
1361
+ if isinstance(components, dict):
1362
+ holders.extend(components.values())
1363
+
1364
+ for holder in holders:
1365
+ if not isinstance(holder, torch.nn.Module):
1366
+ continue
1367
+ for module in holder.modules():
1368
+ if isinstance(module, BaseTunerLayer) and id(module) not in seen:
1369
+ seen.add(id(module))
1370
+ yield module
1371
+
1372
+
1373
+ def load_loras(loras, pipeline):
1374
+ """Load and configure LoRA models."""
1375
+ adapter_names = []
1376
+ adapter_weights = []
1377
+ alphas = {}
1378
+
1379
+ for i, lora in enumerate(loras):
1380
+ model_name = lora.pop("model_name", None)
1381
+ logger.info(f"Loading LoRA: {model_name}")
1382
+ emit_phase("loading", detail=f"LoRA: {model_name}")
1383
+
1384
+ # Use provided adapter_name or generate from index
1385
+ adapter_name = lora.pop("adapter_name", str(i))
1386
+ adapter_names.append(adapter_name)
1387
+
1388
+ # Extract scale for adapter weights - float() because the schema takes a
1389
+ # 'variable:' reference here, and a variable declared as a string default
1390
+ # substitutes as one
1391
+ scale = float(lora.pop("scale", 1.0))
1392
+ adapter_weights.append(scale)
1393
+
1394
+ # Popped before the load: everything left in the dict is a keyword
1395
+ # argument to load_lora_weights, and the alpha is applied to the layers
1396
+ # afterwards rather than passed to it
1397
+ alpha = lora.pop("alpha", None)
1398
+ if alpha is not None:
1399
+ alphas[adapter_name] = alpha
1400
+
1401
+ # Load the LoRA with the adapter name
1402
+ pipeline.load_lora_weights(model_name, adapter_name=adapter_name, **lora)
1403
+
1404
+ for adapter_name, alpha in alphas.items():
1405
+ set_adapter_alpha(pipeline, adapter_name, alpha)
1406
+
1407
+ # Set adapter weights for all loaded LoRAs
1408
+ if adapter_names:
1409
+ logger.info(
1410
+ f"Setting adapter weights: {list(zip(adapter_names, adapter_weights))}"
1411
+ )
1412
+ # Positionally - diffusers' mixin calls the second parameter 'adapter_weights'
1413
+ # while custom pipelines that delegate to the model (ostris/Krea2OstrisEdit)
1414
+ # call it 'weights'
1415
+ pipeline.set_adapters(adapter_names, adapter_weights)
1416
+
1417
+
1418
+ def load_ip_adapter(ip_adapter_definition, pipeline):
1419
+ """Load and configure IP-Adapter if specified."""
1420
+ if ip_adapter_definition is not None:
1421
+ model_name = ip_adapter_definition.pop("model_name")
1422
+ logger.info(f"Loading IP-Adapter: {model_name}")
1423
+ scale = ip_adapter_definition.pop("scale", None)
1424
+ pipeline.load_ip_adapter(model_name, **ip_adapter_definition)
1425
+ if scale is not None:
1426
+ pipeline.set_ip_adapter_scale(scale)
1427
+
1428
+
1429
+ def load_and_configure_scheduler(
1430
+ scheduler_definition, pipeline, component_name="scheduler"
1431
+ ):
1432
+ """Load and configure a pipeline's scheduler if specified.
1433
+
1434
+ A definition does either or both of two things, in that order: replace the
1435
+ scheduler with one built from another type's config, and set the sigma
1436
+ shift on whatever scheduler the pipeline then holds.
1437
+
1438
+ The component is named rather than assumed because a pipeline can carry
1439
+ more than one. MiniMax-H3 steps video and audio latents down two schedules
1440
+ inside a single transformer call - 'scheduler' and 'audio_scheduler', whose
1441
+ shifts (12.0 and 3.0 in the released checkpoint) are set independently, and
1442
+ the video one is what a few-step schedule has to lower: at the checkpoint's
1443
+ 12.0 a five-point sigma grid spends every step above 0.8 and then drops to
1444
+ zero in one, which denoises to noise.
1445
+
1446
+ Args:
1447
+ scheduler_definition: The step's scheduler block, or None
1448
+ pipeline: The loaded pipeline
1449
+ component_name: Which scheduler the definition configures
1450
+ """
1451
+ if scheduler_definition is None:
1452
+ return
1453
+
1454
+ scheduler_configuration = scheduler_definition.get("configuration", None) or {}
1455
+ scheduler_type = scheduler_configuration.get("scheduler_type", None)
1456
+ if scheduler_type is not None:
1457
+ from_config_args = scheduler_definition.get("from_config_args", {})
1458
+ logger.info(f"Loading {component_name}: {scheduler_type}")
1459
+ setattr(
1460
+ pipeline,
1461
+ component_name,
1462
+ scheduler_type.from_config(
1463
+ get_component(pipeline, component_name).config, **from_config_args
1464
+ ),
1465
+ )
1466
+
1467
+ shift = scheduler_definition.get("shift", None)
1468
+ if shift is None:
1469
+ return
1470
+
1471
+ scheduler = get_component(pipeline, component_name)
1472
+ if scheduler is None:
1473
+ raise ValueError(
1474
+ f"Cannot set a shift on '{component_name}' - the pipeline registers "
1475
+ "it but has not loaded it"
1476
+ )
1477
+ if not has_method(scheduler, "set_shift"):
1478
+ raise ValueError(
1479
+ f"{type(scheduler).__name__} does not take a sigma shift - "
1480
+ f"'{component_name}' has no set_shift()"
1481
+ )
1482
+
1483
+ # Instance state the scheduler keeps until its next set_timesteps, which is
1484
+ # the run itself - so this survives loading and every later run of the step
1485
+ logger.info(f"Setting {component_name} shift: {shift}")
1486
+ scheduler.set_shift(float(shift))
1487
+
1488
+
1489
+ def auto_cpu_offload_enabled(configuration):
1490
+ """Whether the configuration asks its components manager to offload to the CPU."""
1491
+ return configuration.get("components_manager", {}).get(
1492
+ "enable_auto_cpu_offload", False
1493
+ )
1494
+
1495
+
1496
+ def auto_cpu_offload_active(configuration, device):
1497
+ """Whether the components manager actually owns device placement.
1498
+
1499
+ Mirrors the MPS skip in create_components_manager() - on MPS the manager
1500
+ never installs its offload hooks, so callers must not assume it owns
1501
+ device placement there.
1502
+ """
1503
+ return auto_cpu_offload_enabled(configuration) and get_device_type(device) != "mps"
1504
+
1505
+
1506
+ def create_components_manager(configuration, device):
1507
+ """Create the components manager for a modular pipeline, when one is configured.
1508
+
1509
+ A ComponentsManager tracks the components of a modular pipeline and can keep only
1510
+ the ones currently running on the device, moving the rest to system memory.
1511
+
1512
+ Args:
1513
+ configuration: Pipeline configuration dictionary
1514
+ device: Device the pipeline runs on
1515
+
1516
+ Returns:
1517
+ A configured ComponentsManager, or None when the pipeline does not use one
1518
+ """
1519
+ manager_configuration = configuration.get("components_manager", None)
1520
+ if manager_configuration is None:
1521
+ return None
1522
+
1523
+ # Imported here because importing modular diffusers warns that it is experimental
1524
+ from diffusers import ComponentsManager
1525
+
1526
+ logger.info("Creating components manager")
1527
+ components_manager = ComponentsManager()
1528
+
1529
+ if auto_cpu_offload_enabled(configuration):
1530
+ # ComponentsManager.enable_auto_cpu_offload() calls device.mem_get_info(),
1531
+ # which torch does not implement for MPS. Unified memory also makes the
1532
+ # feature far less useful there than on CUDA, so skip it rather than fail.
1533
+ if get_device_type(device) == "mps":
1534
+ logger.warning(
1535
+ "components_manager auto CPU offload is not supported on MPS, skipping"
1536
+ )
1537
+ else:
1538
+ offload_arguments = {}
1539
+ memory_reserve_margin = manager_configuration.get(
1540
+ "memory_reserve_margin", None
1541
+ )
1542
+ if memory_reserve_margin is not None:
1543
+ offload_arguments["memory_reserve_margin"] = memory_reserve_margin
1544
+
1545
+ # Enabled before the components load so each one is hooked as it is added
1546
+ logger.info(f"Enabling components manager auto CPU offload on {device}")
1547
+ components_manager.enable_auto_cpu_offload(
1548
+ device=device, **offload_arguments
1549
+ )
1550
+
1551
+ return components_manager
1552
+
1553
+
1554
+ def has_component_group_offload(configuration):
1555
+ """Whether a per-component entry keeps its component off the device.
1556
+
1557
+ The 'components' block is applied by configure_components() after the pipeline is
1558
+ loaded, but load_component() has to decide where to materialize weights and whether
1559
+ to move the pipeline to the device before that block is ever read. A workflow whose
1560
+ only offload configuration lives under components.* still needs both of those
1561
+ earlier decisions to treat it as offloading.
1562
+
1563
+ Group offloading and on-demand residency both qualify: each leaves its component in
1564
+ system memory between uses, so materializing the pipeline on the device first would
1565
+ load in full exactly what these were configured to avoid holding.
1566
+
1567
+ Args:
1568
+ configuration: Configuration of the component being loaded
1569
+
1570
+ Returns:
1571
+ True when any per-component entry keeps its component off the device
1572
+ """
1573
+ components = configuration.get("components") or {}
1574
+ return any(
1575
+ isinstance(settings, dict)
1576
+ and (
1577
+ settings.get("group_offload") is not None
1578
+ or settings.get("residency") == "on_demand"
1579
+ )
1580
+ for settings in components.values()
1581
+ )
1582
+
1583
+
1584
+ def loading_device(configuration):
1585
+ """The device a component's weights are materialized on while it loads.
1586
+
1587
+ Offloading brings each part of a model onto the device only while it runs, so the
1588
+ weights have to land in system memory first. A default torch device pointing at the
1589
+ GPU would build every module directly in VRAM instead, running a large pipeline out
1590
+ of memory before its offload hooks are ever installed.
1591
+
1592
+ Args:
1593
+ configuration: Configuration of the component being loaded
1594
+
1595
+ Returns:
1596
+ A context manager active for the duration of the load
1597
+ """
1598
+ offloads = (
1599
+ configuration.get("offload", None) is not None
1600
+ or configuration.get("group_offload", None) is not None
1601
+ or has_component_group_offload(configuration)
1602
+ )
1603
+
1604
+ if offloads:
1605
+ logger.debug("Loading into system memory - the component will be offloaded")
1606
+ return torch.device("cpu")
1607
+
1608
+ return contextlib.nullcontext()
1609
+
1610
+
1611
+ def get_block_configs(configuration, component):
1612
+ """The block configs a workflow sets on a modular pipeline, checked against it.
1613
+
1614
+ A modular pipeline's blocks declare configs of their own - values they read while
1615
+ they run rather than components or call arguments. MiniMax-H3 declares three, and
1616
+ they are how the canvas the request generates on and the resolution its references
1617
+ are encoded at are set:
1618
+
1619
+ "configs": { "canvas_short_edge": 768, "reference_image_short_edge": 1024 }
1620
+
1621
+ This is deliberately not a knob per config per model. Every modular pipeline
1622
+ declares its own set, `update_components()` sets any of them, and what a workflow
1623
+ may say here is whatever the pipeline it named declares. The names are checked
1624
+ because update_components ignores the ones it does not know with a warning, and a
1625
+ silently dropped config reads as a setting that did nothing.
1626
+
1627
+ Args:
1628
+ configuration: Pipeline configuration dictionary
1629
+ component: The loaded pipeline the configs are for
1630
+
1631
+ Returns:
1632
+ Dict of config name to value, empty when the workflow sets none
1633
+
1634
+ Raises:
1635
+ ValueError: If the pipeline takes no configs, or does not declare one by name
1636
+ """
1637
+ configs = configuration.get("configs", None)
1638
+ if not configs:
1639
+ return {}
1640
+
1641
+ if not has_method(component, "update_components"):
1642
+ raise ValueError(
1643
+ f"'configs' is only supported on modular pipelines, "
1644
+ f"{type(component).__name__} does not have update_components"
1645
+ )
1646
+
1647
+ # The specs a pipeline builds from its blocks. Guarded rather than indexed - a
1648
+ # pipeline that stops keeping them under this name should lose the check, not
1649
+ # the feature
1650
+ declared = getattr(component, "_config_specs", None)
1651
+ if declared is not None:
1652
+ unknown = [name for name in configs if name not in declared]
1653
+ if unknown:
1654
+ raise ValueError(
1655
+ f"{type(component).__name__} declares no config named "
1656
+ f"{', '.join(sorted(unknown))} - the ones it declares are "
1657
+ f"{', '.join(sorted(declared)) or 'none'}"
1658
+ )
1659
+
1660
+ logger.info(f"Setting block configs: {', '.join(configs)}")
1661
+ return dict(configs)
1662
+
1663
+
1664
+ def place_component(
1665
+ component, component_name, configuration, device, components_manager=None
1666
+ ):
1667
+ """Give a loaded component its offloading hooks and its device.
1668
+
1669
+ Split out of load_component because placement has to come last. Every hook
1670
+ here - group offloading, layerwise casting, model or sequential CPU offload -
1671
+ pins the modules and weights that exist when it is installed, so anything that
1672
+ adds or replaces weights afterwards (a LoRA, an IP-Adapter) is left outside the
1673
+ hook's bookkeeping: sequential offload streams the weights it recorded onto the
1674
+ accelerator and the adapter's own tensors are never among them, which runs the
1675
+ step on uninitialized weights and produces NaN.
1676
+
1677
+ Args:
1678
+ component: The loaded pipeline or component
1679
+ component_name: What is being placed, for the log
1680
+ configuration: The component's configuration block
1681
+ device: Device the component runs on
1682
+ components_manager: The modular pipeline's components manager, if it has one
1683
+
1684
+ Returns:
1685
+ The placed component
1686
+ """
1687
+ # Translate before anything reads the backend - the offload downgrades below
1688
+ # have to see the device the component will actually run on
1689
+ device = resolve_device(device)
1690
+
1691
+ # Handle group_offload configuration
1692
+ group_offload_configuration = get_group_offload_configuration(configuration, device)
1693
+ if group_offload_configuration is not None:
1694
+ component.enable_group_offload(**group_offload_configuration)
1695
+
1696
+ # Handle enable_layerwise_casting configuration
1697
+ enable_layerwise_casting_configuration = configuration.get(
1698
+ "enable_layerwise_casting", None
1699
+ )
1700
+ if enable_layerwise_casting_configuration is not None:
1701
+ component.enable_layerwise_casting(**enable_layerwise_casting_configuration)
1702
+
1703
+ # Configure component device settings
1704
+ preserve_device_placement = configuration.get("preserve_device_placement", False)
1705
+ offload = configuration.get("offload", None)
1706
+
1707
+ # Offloading streams a model between system memory and an accelerator - there is
1708
+ # nothing to stream to when the run is on the CPU
1709
+ if offload is not None and get_device_type(device) == "cpu":
1710
+ logger.warning(f"Ignoring '{offload}' offload - {device} is not an accelerator")
1711
+ offload = None
1712
+
1713
+ # Sequential offload streams each submodule onto the accelerator as it runs,
1714
+ # a trade only a separate memory pool rewards. MPS shares one pool with the
1715
+ # CPU, so the streaming hands back no residency and costs a copy per
1716
+ # submodule per step. Model offload keeps the coarse win - idle components
1717
+ # off the Metal allocator - without paying that
1718
+ if offload == "sequential" and get_device_type(device) == "mps":
1719
+ excluded = configuration.get("exclude_from_cpu_offload", [])
1720
+ ignored = (
1721
+ f"; 'exclude_from_cpu_offload' ({', '.join(excluded)}) is sequential-only "
1722
+ "and does not carry over"
1723
+ if excluded
1724
+ else ""
1725
+ )
1726
+ logger.warning(
1727
+ f"Using model offload in place of sequential on {device} - sequential "
1728
+ f"streams weights per submodule, which buys back no memory on unified "
1729
+ f"memory{ignored}"
1730
+ )
1731
+ offload = "model"
1732
+
1733
+ if offload == "model":
1734
+ logger.debug(f"Enabling model CPU offload onto {device}")
1735
+ component.enable_model_cpu_offload(device=device)
1736
+ elif offload == "sequential":
1737
+ logger.debug(f"Enabling sequential CPU offload onto {device}")
1738
+ for excluded_name in configuration.get("exclude_from_cpu_offload", []):
1739
+ logger.debug(f"Excluding {excluded_name} from CPU offload")
1740
+ component._exclude_from_cpu_offload.append(excluded_name)
1741
+ component.enable_sequential_cpu_offload(device=device)
1742
+ elif components_manager is not None and auto_cpu_offload_active(
1743
+ configuration, device
1744
+ ):
1745
+ # Moving everything to the device here would defeat the offloading - the
1746
+ # manager's hooks bring each component on device as the pipeline needs it
1747
+ logger.debug("Device placement is owned by the components manager")
1748
+ elif has_component_group_offload(configuration):
1749
+ # configure_components() installs group-offload hooks per-component after
1750
+ # this returns - moving the whole pipeline to the device now would load it
1751
+ # in full before those hooks exist, defeating the offloading
1752
+ logger.info(
1753
+ f"components configure group offloading - not moving pipeline to {device}"
1754
+ )
1755
+ elif hasattr(component, "to") and not preserve_device_placement:
1756
+ logger.debug(f"Moving {component_name} to device: {device}")
1757
+ component = component.to(device)
1758
+
1759
+ return component
1760
+
1761
+
1762
+ def load_component(
1763
+ component_name,
1764
+ configuration,
1765
+ from_pretrained_arguments,
1766
+ device,
1767
+ reused_components=None,
1768
+ defer_placement=False,
1769
+ ):
1770
+ """Load and configure a pipeline or component.
1771
+
1772
+ Args:
1773
+ component_name: What is being loaded, for the log
1774
+ configuration: The component's configuration block
1775
+ from_pretrained_arguments: Arguments for the constructor
1776
+ device: Device the component is loaded for
1777
+ reused_components: Components an earlier step shared into this one, by name
1778
+ defer_placement: Load the component without placing it - the caller calls
1779
+ place_component once it has finished altering the weights
1780
+ """
1781
+ component_type = configuration["component_type"]
1782
+ component = None
1783
+
1784
+ # Refused before anything reaches the Hub: an untrusted workflow must not
1785
+ # be able to have diffusers fetch and run code on its behalf
1786
+ require_trusted_from_pretrained_arguments(from_pretrained_arguments, component_name)
1787
+
1788
+ # A standard pipeline takes a component as a constructor argument. A modular one
1789
+ # cannot: it is built from the component specs in its own index and given the
1790
+ # objects afterwards, which is also what keeps load_components() from pulling a
1791
+ # second copy of the weights - it skips the components already registered
1792
+ reused_components = reused_components or {}
1793
+ takes_components_after_load = has_method(component_type, "update_components")
1794
+ if reused_components and not takes_components_after_load:
1795
+ from_pretrained_arguments.update(reused_components)
1796
+
1797
+ # A modular pipeline can hand its components to a ComponentsManager, which then
1798
+ # owns their device placement
1799
+ components_manager = create_components_manager(configuration, device)
1800
+ if components_manager is not None:
1801
+ from_pretrained_arguments["components_manager"] = components_manager
1802
+
1803
+ # MPS (Apple Silicon) has numerical instability with float16 matmul operations,
1804
+ # producing NaN values that result in black images. The dtype is left as asked for -
1805
+ # silently loading a model in a dtype the workflow did not request would be worse -
1806
+ # so this only warns.
1807
+ if (
1808
+ get_device_type(device) == "mps"
1809
+ and from_pretrained_arguments.get("torch_dtype") == torch.float16
1810
+ ):
1811
+ logger.warning(
1812
+ f"On MPS devices float16 produces NaN values on Apple Silicon"
1813
+ f"Consider changing torch_dtype from float16 to float32 for {component_name} "
1814
+ )
1815
+
1816
+ model_name = None
1817
+ try:
1818
+ with loading_device(configuration):
1819
+ # Load from model name
1820
+ if "model_name" in from_pretrained_arguments:
1821
+ model_name = from_pretrained_arguments.pop("model_name")
1822
+ logger.info(f"Loading {component_name} from model: {model_name}")
1823
+ emit_phase("loading", detail=f"{component_name}: {model_name}")
1824
+ with download_watch.watch(
1825
+ model_name, cache_dir=from_pretrained_arguments.get("cache_dir")
1826
+ ):
1827
+ component = component_type.from_pretrained(
1828
+ model_name, **from_pretrained_arguments
1829
+ )
1830
+
1831
+ # Load from single file
1832
+ elif "from_single_file" in from_pretrained_arguments:
1833
+ from_single_file = from_pretrained_arguments.pop("from_single_file")
1834
+ logger.info(
1835
+ f"Loading {component_name} from single file: {from_single_file}"
1836
+ )
1837
+ emit_phase("loading", detail=f"{component_name}: {from_single_file}")
1838
+ component = component_type.from_single_file(
1839
+ from_single_file, **from_pretrained_arguments
1840
+ )
1841
+
1842
+ # Create new component
1843
+ else:
1844
+ logger.info(f"Creating new {component_name}")
1845
+ component = component_type(**from_pretrained_arguments)
1846
+
1847
+ # Register the shared components before anything is pulled, so the
1848
+ # weights an earlier step already loaded and quantized are the ones
1849
+ # this step runs on rather than a second copy of them. The block
1850
+ # configs go in the same call - update_components takes both
1851
+ update_arguments = get_block_configs(configuration, component)
1852
+ if reused_components and takes_components_after_load:
1853
+ logger.info(
1854
+ f"Reusing {', '.join(reused_components)} from an earlier step"
1855
+ )
1856
+ update_arguments.update(reused_components)
1857
+ if update_arguments:
1858
+ component.update_components(**update_arguments)
1859
+
1860
+ # Modular pipelines load only their config in from_pretrained - the component
1861
+ # weights are pulled separately by load_components()
1862
+ load_components_arguments = get_load_components_arguments(configuration)
1863
+ if load_components_arguments is not None:
1864
+ if not has_method(component, "load_components"):
1865
+ raise ValueError(
1866
+ f"load_components is only supported on modular pipelines, "
1867
+ f"{component_type.__name__} does not have it"
1868
+ )
1869
+ logger.info(f"Loading components for {component_name}")
1870
+ component.load_components(**load_components_arguments)
1871
+
1872
+ if defer_placement:
1873
+ # The caller places this itself, once it has finished loading the
1874
+ # things that alter the weights - see place_component
1875
+ logger.debug(f"Deferring placement of {component_name}")
1876
+ return component
1877
+
1878
+ return place_component(
1879
+ component, component_name, configuration, device, components_manager
1880
+ )
1881
+
1882
+ except WorkflowCancelled:
1883
+ # A cancel that aborted a download (dw/download_watch.py) - not a
1884
+ # load failure, so no error log
1885
+ raise
1886
+ except Exception as e:
1887
+ # 401/403 from the Hub means the account behind whatever token (or
1888
+ # lack of one) HfApi is using cannot read this repo - almost always
1889
+ # a gated model the user has not requested access to, or has not
1890
+ # logged in for. A whole-pipeline load raises the HfHubHTTPError
1891
+ # itself; a per-component load gets it wrapped in an EnvironmentError
1892
+ # by diffusers' _get_model_file / transformers' cached_file, so the
1893
+ # cause chain is searched. Every other error (a real outage, a bad
1894
+ # repo id) is logged once with its traceback and re-raised unchanged
1895
+ if _hub_auth_status(e) is not None:
1896
+ repo = model_name or component_name
1897
+ logger.error(
1898
+ f"Hugging Face authentication required loading {component_name} "
1899
+ f"({repo}): {e}",
1900
+ exc_info=True,
1901
+ )
1902
+ raise RuntimeError(
1903
+ f"Model '{repo}' requires Hugging Face authentication - run "
1904
+ f"'huggingface-cli login', or request access at "
1905
+ f"https://huggingface.co/{repo}"
1906
+ ) from e
1907
+ logger.error(f"{type(e).__name__} loading {component_name}: {e}", exc_info=True)
1908
+ raise
1909
+
1910
+
1911
+ def _hub_auth_status(error):
1912
+ """The 401/403 status an exception (or anything in its cause/context
1913
+ chain) carries from the Hub, else None."""
1914
+ seen = set()
1915
+ while error is not None and id(error) not in seen:
1916
+ seen.add(id(error))
1917
+ if isinstance(error, HfHubHTTPError):
1918
+ status = getattr(getattr(error, "response", None), "status_code", None)
1919
+ if status in (401, 403):
1920
+ return status
1921
+ error = error.__cause__ or error.__context__
1922
+ return None
1923
+
1924
+
1925
+ def apply_sdnq_optimizations(pipeline, component_names):
1926
+ """Apply SDNQ quantized matmul optimization to pipeline components.
1927
+
1928
+ Uses sdnq's apply_sdnq_options_to_model to enable INT8 matmul
1929
+ on supported hardware (CUDA, XPU).
1930
+
1931
+ Args:
1932
+ pipeline: The loaded diffusers pipeline
1933
+ component_names: List of component names to optimize (e.g., ["transformer", "text_encoder"])
1934
+ """
1935
+ try:
1936
+ from sdnq.loader import apply_sdnq_options_to_model
1937
+ from sdnq.common import use_torch_compile as triton_is_available
1938
+ except ImportError:
1939
+ logger.warning("sdnq not installed, skipping SDNQ optimizations")
1940
+ return
1941
+
1942
+ if not triton_is_available:
1943
+ logger.info("Triton not available, skipping SDNQ quantized matmul optimization")
1944
+ return
1945
+
1946
+ if not (
1947
+ torch.cuda.is_available() or hasattr(torch, "xpu") and torch.xpu.is_available()
1948
+ ):
1949
+ logger.info(
1950
+ "SDNQ quantized matmul requires CUDA or XPU, skipping on this device"
1951
+ )
1952
+ return
1953
+
1954
+ for name in component_names:
1955
+ # A missing name (typo) and a registered-but-unloaded one both mean "nothing
1956
+ # to optimize here" for this call - same warn-and-skip either way
1957
+ try:
1958
+ component = get_component(pipeline, name)
1959
+ except ValueError:
1960
+ component = None
1961
+
1962
+ if component is not None:
1963
+ logger.info(f"Applying SDNQ quantized matmul to {name}")
1964
+ setattr(
1965
+ pipeline,
1966
+ name,
1967
+ apply_sdnq_options_to_model(component, use_quantized_matmul=True),
1968
+ )
1969
+ else:
1970
+ logger.warning(
1971
+ f"Component '{name}' not found on pipeline, skipping SDNQ optimization"
1972
+ )
1973
+
1974
+
1975
+ def get_cache_transformer(pipeline):
1976
+ """Find the denoiser a cache hook attaches to.
1977
+
1978
+ Most pipelines register theirs as 'transformer', but a modular pipeline names
1979
+ it after the workflow it serves - MiniMax-H3's ref2va denoises through
1980
+ 'transformer_ref'. Looking only for 'transformer' silently skips caching on
1981
+ those, so try the alternates diffusers' modular pipelines actually use.
1982
+
1983
+ Args:
1984
+ pipeline: The loaded diffusers pipeline
1985
+
1986
+ Returns:
1987
+ The transformer component, or None when the pipeline has none
1988
+ """
1989
+ for name in ("transformer", "transformer_ref"):
1990
+ transformer = getattr(pipeline, name, None)
1991
+ if transformer is not None:
1992
+ return transformer
1993
+ return None
1994
+
1995
+
1996
+ class _ReportingProgressBar:
1997
+ """A tqdm bar that also reports each advance to the active run.
1998
+
1999
+ Wraps rather than subclasses, because the bar it wraps is whatever the
2000
+ block's own progress_bar() built - tqdm, or a notebook bar, or whatever
2001
+ a future diffusers uses. Everything it does not intercept falls through
2002
+ to the real bar, so the terminal output is unchanged.
2003
+ """
2004
+
2005
+ def __init__(self, bar, on_advance, total=None):
2006
+ self._bar = bar
2007
+ self._on_advance = on_advance
2008
+ # Counted here rather than read off the bar: a disabled tqdm - which
2009
+ # is what a quiet server or a notebook config leaves you with - keeps
2010
+ # its own `n` at zero while still being advanced normally
2011
+ self._done = 0
2012
+ self._total = total if total is not None else getattr(bar, "total", None)
2013
+
2014
+ def update(self, n=1):
2015
+ result = self._bar.update(n)
2016
+ self._done += n or 0
2017
+ self._on_advance(self._done, self._total)
2018
+ return result
2019
+
2020
+ def __iter__(self):
2021
+ # Reported after the body of the loop has run, not before it: the
2022
+ # step is finished when control comes back here
2023
+ for item in self._bar:
2024
+ yield item
2025
+ self._done += 1
2026
+ self._on_advance(self._done, self._total)
2027
+
2028
+ def __enter__(self):
2029
+ self._bar.__enter__()
2030
+ return self
2031
+
2032
+ def __exit__(self, *exception):
2033
+ return self._bar.__exit__(*exception)
2034
+
2035
+ def __getattr__(self, name):
2036
+ # Guarded: the wrapped bar is the first thing __init__ sets, and an
2037
+ # unguarded lookup of it before then recurses forever
2038
+ if name == "_bar":
2039
+ raise AttributeError(name)
2040
+ return getattr(self._bar, name)
2041
+
2042
+
2043
+ def _progress_bar_holders(pipeline):
2044
+ """Every object under a modular pipeline that can open a progress bar.
2045
+
2046
+ The denoise loop is a block, not the pipeline, and it calls its own
2047
+ `self.progress_bar(...)` - so the tree is what has to be walked. Uses
2048
+ `_blocks`, not the public `blocks`, which hands back a deepcopy: patching
2049
+ a copy would report nothing and look like this never worked.
2050
+ """
2051
+ holders = []
2052
+ seen = set()
2053
+
2054
+ def walk(candidate):
2055
+ if candidate is None or id(candidate) in seen:
2056
+ return
2057
+ seen.add(id(candidate))
2058
+ # __dict__, because the patch is an instance attribute: an object
2059
+ # with none could not be patched and must not be tried
2060
+ if callable(getattr(candidate, "progress_bar", None)) and hasattr(
2061
+ candidate, "__dict__"
2062
+ ):
2063
+ holders.append(candidate)
2064
+ children = getattr(candidate, "sub_blocks", None)
2065
+ if hasattr(children, "values"):
2066
+ for child in children.values():
2067
+ walk(child)
2068
+
2069
+ walk(pipeline)
2070
+ walk(getattr(pipeline, "_blocks", None))
2071
+ return holders
2072
+
2073
+
2074
+ @contextlib.contextmanager
2075
+ def reported_progress_bars(pipeline):
2076
+ """Report each denoise step of a pipeline that takes no step callback.
2077
+
2078
+ A ModularPipeline - H3, LTX-2, Qwen-Image and every family diffusers has
2079
+ moved over - has no `callback_on_step_end` parameter, so the whole
2080
+ denoise loop passed in silence: one 'generating' phase, then nothing for
2081
+ however many minutes it took, which reads exactly like a hung run. What
2082
+ those blocks do have is a tqdm bar, and every advance of it is a step.
2083
+
2084
+ The patch is per-instance and undone on the way out, so a pipeline this
2085
+ process keeps loaded is handed back as it was found.
2086
+ """
2087
+ holders = _progress_bar_holders(pipeline)
2088
+ if not holders:
2089
+ yield
2090
+ return
2091
+
2092
+ run_context = get_context()
2093
+
2094
+ def on_advance(done, total):
2095
+ run_context.emit("pipeline_step", step=done, total_steps=total)
2096
+ # Past the last step there is still the decode, which on video is
2097
+ # minutes with the bar sitting at 100%
2098
+ if done is not None and total is not None and done >= total:
2099
+ emit_phase("decoding")
2100
+ # The one cancellation checkpoint inside a modular denoise loop:
2101
+ # without it a cancel waits out the whole generation
2102
+ run_context.check_cancelled()
2103
+
2104
+ patched = []
2105
+ for holder in holders:
2106
+ original = holder.progress_bar
2107
+ # Whether the name was already an attribute of the instance decides
2108
+ # how it is put back: restored, or removed so the class method shows
2109
+ # through again rather than a bound copy of it being frozen on
2110
+ patched.append((holder, original, "progress_bar" in vars(holder)))
2111
+
2112
+ def reporting(iterable=None, total=None, _original=original):
2113
+ bar = _original(iterable=iterable, total=total)
2114
+ if total is None and iterable is not None:
2115
+ # An iterated bar's total is the length of what it iterates,
2116
+ # when that can be known at all
2117
+ total = getattr(bar, "total", None)
2118
+ return _ReportingProgressBar(bar, on_advance, total)
2119
+
2120
+ holder.progress_bar = reporting
2121
+ try:
2122
+ yield
2123
+ finally:
2124
+ for holder, original, was_own in patched:
2125
+ if was_own:
2126
+ holder.progress_bar = original
2127
+ else:
2128
+ try:
2129
+ del holder.progress_bar
2130
+ except AttributeError:
2131
+ holder.progress_bar = original
2132
+
2133
+
2134
+ def _runs_its_blocks_in_sequence(blocks):
2135
+ """Whether a modular block container runs every sub-block in order.
2136
+
2137
+ diffusers' own `SequentialPipelineBlocks` is the answer; the import is
2138
+ local and forgiving because a pipeline that is not modular at all never
2139
+ reaches here, and a diffusers without the class is one with no modular
2140
+ pipelines to narrate.
2141
+ """
2142
+ try:
2143
+ from diffusers.modular_pipelines.modular_pipeline import (
2144
+ SequentialPipelineBlocks,
2145
+ )
2146
+ except ImportError: # pragma: no cover - a diffusers without modular
2147
+ return False
2148
+ return isinstance(blocks, SequentialPipelineBlocks)
2149
+
2150
+
2151
+ @contextlib.contextmanager
2152
+ def reported_blocks(pipeline, label):
2153
+ """Name each of a modular pipeline's top-level blocks as it starts.
2154
+
2155
+ The denoise loop is only one of them, and on a reference-conditioned
2156
+ model it is not the long one: encoding a video reference runs for
2157
+ minutes inside `vae_encoder` before a single bar advances, so the whole
2158
+ lead-in went by with nothing emitted and a consumer could not tell it
2159
+ from a hang (#95). The blocks are named - `before_encode`,
2160
+ `text_encoder`, `vae_encoder`, `denoise`, `decode` on H3 - and naming
2161
+ each one as it begins turns that silence into "it is encoding the
2162
+ reference", plus a `seconds_since_event` that resets at every boundary.
2163
+
2164
+ Reported as `log` events rather than phases: `PHASES` is a closed set a
2165
+ consumer switches on, and a block name is a detail, not a new state.
2166
+
2167
+ The patch is on the class because `block(pipeline, state)` resolves
2168
+ `__call__` on the type, not the instance - so it is guarded by identity
2169
+ (only the pipeline's own top-level blocks report) and undone on the way
2170
+ out.
2171
+ """
2172
+ blocks = getattr(pipeline, "_blocks", None)
2173
+ sub_blocks = getattr(blocks, "sub_blocks", None)
2174
+ if blocks is None or not hasattr(sub_blocks, "items"):
2175
+ yield
2176
+ return
2177
+ # Only a sequence runs all of its sub-blocks. A conditional container
2178
+ # (AutoPipelineBlocks, which is what several of H3's own steps are)
2179
+ # *picks* one on its inputs, so narrating it by walking the mapping
2180
+ # would run every branch - hence the check for the one dispatch this
2181
+ # reproduces rather than a duck-typed `sub_blocks`
2182
+ if not _runs_its_blocks_in_sequence(blocks):
2183
+ yield
2184
+ return
2185
+
2186
+ holder = type(blocks)
2187
+ original = holder.__call__
2188
+ was_own = "__call__" in vars(holder)
2189
+ run_context = get_context()
2190
+
2191
+ @torch.no_grad()
2192
+ def reporting(self, pipe, state):
2193
+ # A nested SequentialPipelineBlocks shares the class; only the
2194
+ # pipeline's own top-level sequence is the one worth narrating
2195
+ if self is not blocks:
2196
+ return original(self, pipe, state)
2197
+ for name, block in self.sub_blocks.items():
2198
+ run_context.emit("log", message=f"{label}: {name}")
2199
+ # A block boundary is a cancellation checkpoint the lead-in
2200
+ # otherwise has none of
2201
+ run_context.check_cancelled()
2202
+ try:
2203
+ pipe, state = block(pipe, state)
2204
+ except WorkflowCancelled:
2205
+ raise
2206
+ except Exception:
2207
+ # What diffusers' own dispatch logs, kept because this
2208
+ # replaces that loop
2209
+ logger.error(f"Error in block: ({name}, {block.__class__.__name__})")
2210
+ raise
2211
+ return pipe, state
2212
+
2213
+ holder.__call__ = reporting
2214
+ try:
2215
+ yield
2216
+ finally:
2217
+ if was_own:
2218
+ holder.__call__ = original
2219
+ else:
2220
+ try:
2221
+ del holder.__call__
2222
+ except AttributeError:
2223
+ holder.__call__ = original
2224
+
2225
+
2226
+ @contextlib.contextmanager
2227
+ def stateful_cache_context(pipeline):
2228
+ """Provide the context a stateful cache hook reads its state through.
2229
+
2230
+ first_block, mag and layer_skip keep per-context state, and their hooks go
2231
+ through diffusers' StateManager, which raises "No context is set" unless a
2232
+ context is active. A DiffusionPipeline sets one around each denoising step and
2233
+ clears the state afterwards in maybe_free_model_hooks; ModularPipeline is not a
2234
+ DiffusionPipeline and does neither, so caching a modular pipeline dies on the
2235
+ first step - and would otherwise carry the previous run's residuals into the
2236
+ next run of a pipeline this process keeps loaded.
2237
+
2238
+ One context spans the whole call rather than each step. The state is keyed by
2239
+ context name, so re-entering per step only re-reads the same entry. Pipelines
2240
+ that run separate conditional and unconditional passes name a context per pass
2241
+ to keep their caches apart, which a shared context would defeat - but a modular
2242
+ pipeline that needed that would be setting its own contexts already, and this
2243
+ is a no-op for pipelines whose cache is not enabled.
2244
+ """
2245
+ transformer = get_cache_transformer(pipeline)
2246
+ if transformer is None or not getattr(transformer, "is_cache_enabled", False):
2247
+ yield
2248
+ return
2249
+
2250
+ logger.debug(f"Entering cache context for {transformer.__class__.__name__}")
2251
+ try:
2252
+ with transformer.cache_context(_CACHE_CONTEXT_NAME):
2253
+ yield
2254
+ finally:
2255
+ # Private, but it is what diffusers' own pipelines call and there is no
2256
+ # public equivalent. Also clears the context an errored call left set
2257
+ transformer._reset_stateful_cache()
2258
+
2259
+
2260
+ def enable_cache_on_transformer(pipeline, cache_config):
2261
+ """Enable cache configuration on the pipeline's transformer.
2262
+
2263
+ Args:
2264
+ pipeline: The loaded diffusers pipeline
2265
+ cache_config: Cache configuration object from get_cache_configuration()
2266
+ """
2267
+ transformer = get_cache_transformer(pipeline)
2268
+ if transformer is None:
2269
+ logger.warning("Pipeline has no transformer, skipping cache configuration")
2270
+ return
2271
+
2272
+ if not hasattr(transformer, "enable_cache"):
2273
+ logger.warning(
2274
+ f"{transformer.__class__.__name__} does not support enable_cache(), skipping"
2275
+ )
2276
+ return
2277
+
2278
+ # FasterCache decides skipping from the pipeline's current timestep. The
2279
+ # callback is a callable, which workflow JSON cannot express, and diffusers
2280
+ # calls it unconditionally on every denoiser forward - left None, the first
2281
+ # inference step dies. Wire it to the pipeline here, where both exist
2282
+ if (
2283
+ cache_config.__class__.__name__ == "FasterCacheConfig"
2284
+ and getattr(cache_config, "current_timestep_callback", None) is None
2285
+ ):
2286
+ logger.debug("Wiring FasterCache current_timestep_callback to the pipeline")
2287
+ cache_config.current_timestep_callback = lambda: pipeline._current_timestep
2288
+
2289
+ # first_block, mag and layer_skip resolve the transformer's block class
2290
+ # through diffusers' registry and raise when it is absent - fill in the
2291
+ # blocks diffusers has not registered before handing the config over
2292
+ register_cache_blocks()
2293
+
2294
+ logger.info(
2295
+ f"Enabling {cache_config.__class__.__name__} on {transformer.__class__.__name__}"
2296
+ )
2297
+ transformer.enable_cache(cache_config)