multimodalart HF Staff commited on
Commit
186aa49
·
verified ·
1 Parent(s): d11a565

Sync the split MiniMax-H3 Spaces

Browse files
This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. README.md +91 -7
  2. app.py +306 -0
  3. diffusers/__init__.py +1750 -0
  4. diffusers/callbacks.py +244 -0
  5. diffusers/commands/__init__.py +27 -0
  6. diffusers/commands/custom_blocks.py +140 -0
  7. diffusers/commands/diffusers_cli.py +69 -0
  8. diffusers/commands/env.py +185 -0
  9. diffusers/commands/fp16_safetensors.py +144 -0
  10. diffusers/commands/run.py +1227 -0
  11. diffusers/commands/schema.py +287 -0
  12. diffusers/commands/skills.py +344 -0
  13. diffusers/configuration_utils.py +752 -0
  14. diffusers/dependency_versions_check.py +34 -0
  15. diffusers/dependency_versions_table.py +57 -0
  16. diffusers/experimental/README.md +5 -0
  17. diffusers/experimental/__init__.py +1 -0
  18. diffusers/experimental/rl/__init__.py +1 -0
  19. diffusers/experimental/rl/value_guided_sampling.py +153 -0
  20. diffusers/guiders/__init__.py +31 -0
  21. diffusers/guiders/adaptive_projected_guidance.py +253 -0
  22. diffusers/guiders/adaptive_projected_guidance_mix.py +297 -0
  23. diffusers/guiders/auto_guidance.py +198 -0
  24. diffusers/guiders/classifier_free_guidance.py +156 -0
  25. diffusers/guiders/classifier_free_zero_star_guidance.py +164 -0
  26. diffusers/guiders/frequency_decoupled_guidance.py +335 -0
  27. diffusers/guiders/guider_utils.py +396 -0
  28. diffusers/guiders/magnitude_aware_guidance.py +159 -0
  29. diffusers/guiders/perturbed_attention_guidance.py +289 -0
  30. diffusers/guiders/skip_layer_guidance.py +280 -0
  31. diffusers/guiders/smoothed_energy_guidance.py +269 -0
  32. diffusers/guiders/tangential_classifier_free_guidance.py +151 -0
  33. diffusers/hooks/__init__.py +30 -0
  34. diffusers/hooks/_common.py +61 -0
  35. diffusers/hooks/_helpers.py +401 -0
  36. diffusers/hooks/context_parallel.py +382 -0
  37. diffusers/hooks/faster_cache.py +654 -0
  38. diffusers/hooks/first_block_cache.py +258 -0
  39. diffusers/hooks/group_offloading.py +1056 -0
  40. diffusers/hooks/hooks.py +312 -0
  41. diffusers/hooks/layer_skip.py +263 -0
  42. diffusers/hooks/layerwise_casting.py +240 -0
  43. diffusers/hooks/mag_cache.py +468 -0
  44. diffusers/hooks/pyramid_attention_broadcast.py +314 -0
  45. diffusers/hooks/smoothed_energy_guidance_utils.py +166 -0
  46. diffusers/hooks/taylorseer_cache.py +345 -0
  47. diffusers/hooks/text_kv_cache.py +173 -0
  48. diffusers/hooks/utils.py +43 -0
  49. diffusers/image_processor.py +1468 -0
  50. diffusers/loaders/__init__.py +159 -0
README.md CHANGED
@@ -1,13 +1,97 @@
1
  ---
2
- title: Minimax H3 Generator
3
- emoji: 📉
4
- colorFrom: pink
5
- colorTo: purple
6
  sdk: gradio
7
- sdk_version: 6.22.0
8
- python_version: '3.13'
9
  app_file: app.py
10
  pinned: false
 
 
11
  ---
12
 
13
- Check out the configuration reference at https://huggingface.co/docs/hub/spaces-config-reference
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  ---
2
+ title: MiniMax-H3 (split, bf16)
3
+ emoji: 🎬
4
+ colorFrom: purple
5
+ colorTo: indigo
6
  sdk: gradio
7
+ sdk_version: 6.20.0
 
8
  app_file: app.py
9
  pinned: false
10
+ short_description: Unquantized MiniMax-H3, split across two ZeroGPU Spaces
11
+ suggested_hardware: zero-a10g
12
  ---
13
 
14
+ # MiniMax-H3 unquantized, split across two Spaces
15
+
16
+ Joint video **and** soundtrack out of a single denoising pass, at **bfloat16 with no quantization anywhere**.
17
+
18
+ This Space is the denoising half: the 61.73 GiB transformer and the two autoencoders. The 62.14 GiB Qwen3-VL
19
+ conditioner runs in
20
+ [`minimax-h3-conditioner`](https://huggingface.co/spaces/diffusers-internal-dev/minimax-h3-conditioner), which this
21
+ Space calls over the gradio API for every request.
22
+
23
+ ## Why split
24
+
25
+ MiniMax-H3 is 195.9 GiB in bfloat16 and a ZeroGPU Space is evicted at **150 GB of storage**. An unquantized single
26
+ Space is therefore impossible — the existing demos
27
+ ([`minimax-h3`](https://huggingface.co/spaces/diffusers-internal-dev/minimax-h3),
28
+ [`-fp8`](https://huggingface.co/spaces/diffusers-internal-dev/minimax-h3-fp8)) run NVFP4 and float8 weights for that
29
+ reason alone. Cut the `MiniMaxH3Blocks` sequence at its `text_encoder` step and both halves fit unquantized:
30
+
31
+ | Space | Subfolders | Download | Resident |
32
+ |---|---|---|---|
33
+ | [`minimax-h3-conditioner`](https://huggingface.co/spaces/diffusers-internal-dev/minimax-h3-conditioner) | `text_encoder/` + `tokenizer/` + `processor/` | 66.7 GB | 62.15 GiB bf16 |
34
+ | this one | `transformer/` + `vae/` + `audio_vae/` | 77.3 GB | 61.73 GiB bf16 + ~20.5 GiB float32 |
35
+
36
+ Besides the quality argument, unquantized weights are the ones AoTI can export; an NVFP4 checkpoint cannot be
37
+ exported at all.
38
+
39
+ ## How the split is expressed
40
+
41
+ `MiniMaxH3Blocks` is a `SequentialPipelineBlocks` of eight steps:
42
+
43
+ ```
44
+ setup -> text_encoder -> vae_encoder -> prepare_layout -> prepare_latents -> set_timesteps -> denoise -> decode
45
+ ```
46
+
47
+ `h3_split_blocks.py` subclasses it with the `text_encoder` step removed. Dropping the step drops the three
48
+ components it declares, so `load_components` resolves `transformer` / `vae` / `audio_vae` / the two schedulers out of
49
+ the shared `modular_model_index.json` and never fetches the conditioner — and `prompt_embeds` and `text_token_tags`
50
+ become ordinary required inputs of the pipeline call:
51
+
52
+ ```py
53
+ pipe = MiniMaxH3GeneratorBlocks().init_pipeline("diffusers-internal-dev/MiniMax-H3")
54
+ pipe.load_components(dtype=torch.bfloat16)
55
+ state = pipe(prompt_embeds=..., text_token_tags=..., height=768, width=1344, num_frames=124, num_inference_steps=30)
56
+ ```
57
+
58
+ The wire format is exactly those two tensors — `(1, num_text_tokens, 5120)` bfloat16 and `(num_text_tokens,)` int64 —
59
+ carried as one safetensors file with the resolved `height` / `width` / `num_frames` in its metadata header. A
60
+ text-only request is 246 KB of it; one 768x1344 keyframe adds 1016 vision rows and takes it to 10.7 MB.
61
+
62
+ The `setup` step runs on **both** halves. It owns no component (PIL and arithmetic) and it resolves the canvas, the
63
+ `17 * n + 5` frame count and the keyframes placed onto that canvas — which the conditioner needs to build its vision
64
+ blocks and this Space needs to encode with the video VAE. It is deterministic, and the conditioner returns the plan
65
+ it resolved so this Space pins the same canvas rather than re-deriving it.
66
+
67
+ ## Nothing is paid for with GPU time
68
+
69
+ Download, load and the move onto the card happen at **startup**: `import spaces` at module top patches `torch.cuda`
70
+ before any GPU is attached, and these are plain bfloat16 tensors, so ZeroGPU's startup packing handles them (it is a
71
+ torchao `Float8Tensor` that cannot be packed, and there is none here). The conditioner round trip is a network call
72
+ on this Space's CPU. A `@spaces.GPU` call is therefore only the denoise loop and the two decoders.
73
+
74
+ ## Generation constraints
75
+
76
+ Fixed by the checkpoint: 24 fps, a 768 pixel short edge, 5 to 15 s, `num_frames` snapped up to the next `17 * n + 5`,
77
+ no CFG and no negative prompt (it is guidance-distilled, so every step is one forward pass).
78
+
79
+ ## Space variables
80
+
81
+ | Variable | Default | Meaning |
82
+ |---|---|---|
83
+ | `H3_CONDITIONER` | `diffusers-internal-dev/minimax-h3-conditioner` | The Space this one asks for embeddings. |
84
+ | `H3_PLACEMENT` | `resident` | `resident` keeps all 82.3 GiB of weights on the 95.0 GiB card; `offload` hands placement to `ComponentsManager.enable_auto_cpu_offload`. |
85
+ | `H3_ATTENTION` | `_native_cudnn` | cuDNN's fused kernel, 10–20% faster than the SDPA default and needs nothing installed. flash-attention 3 is sm90-only and this pool is sm120. |
86
+ | `H3_GPU_DURATION` | `900` | Seconds per request; the pool applies a 1.5 duration factor. |
87
+ | `H3_GPU_SIZE` | `xlarge` | ZeroGPU allocation size. `large` does not fit. |
88
+
89
+ ## Required secret
90
+
91
+ `HF_TOKEN` — `diffusers-internal-dev/MiniMax-H3` is private, and so is the conditioner Space this one calls.
92
+
93
+ ## Where diffusers comes from
94
+
95
+ MiniMax-H3 is modular-only and not in a released `diffusers`, so the integration branch's `src/diffusers` tree is
96
+ vendored here as a top-level `diffusers/` package; the working directory comes first on `sys.path`, so there is no
97
+ install step. `requirements.txt` only carries what that tree imports.
app.py ADDED
@@ -0,0 +1,306 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """MiniMax-H3, split deployment — **the denoising half**.
2
+
3
+ This Space holds the transformer and the two autoencoders, **unquantized bfloat16**, and nothing else. The 62.14 GiB
4
+ Qwen3-VL conditioner lives in its own Space,
5
+ [`minimax-h3-conditioner`](https://huggingface.co/spaces/diffusers-internal-dev/minimax-h3-conditioner), which this
6
+ one calls over the gradio API for every request; what comes back is a safetensors file holding the two tensors the
7
+ denoiser needs, `prompt_embeds` and `text_token_tags`.
8
+
9
+ Why split at all: MiniMax-H3 is 195.9 GiB in bfloat16 and a ZeroGPU Space is evicted at 150 GB of storage, so an
10
+ unquantized single Space is impossible — the existing demos run NVFP4 or float8 weights for that reason alone. Cut at
11
+ the text-encoder step, this half pulls 77.3 GB (`transformer/` 61.73 GiB + `vae/` 9.70 + `audio_vae/` 0.56) and the
12
+ other 66.7 GB, and neither is quantized.
13
+
14
+ The blockset is `MiniMaxH3Blocks` with its `text_encoder` step removed — see `h3_split_blocks.py`. Dropping the step
15
+ drops the three components it declares, so `load_components` never fetches the conditioner, and `prompt_embeds` /
16
+ `text_token_tags` become ordinary required inputs of the pipeline call.
17
+ """
18
+
19
+ from __future__ import annotations
20
+
21
+ import os
22
+ import tempfile
23
+ import time
24
+ import traceback
25
+
26
+ # First, and at module level. `import spaces` patches `torch.cuda` before any GPU is attached, which is what lets the
27
+ # 82 GiB load happen at **startup** rather than on GPU time; it also has to precede anything that initializes CUDA.
28
+ import spaces
29
+ import gradio as gr
30
+
31
+ MODEL_REPO = os.environ.get("H3_MODEL_REPO", "diffusers-internal-dev/MiniMax-H3")
32
+ CONDITIONER_SPACE = os.environ.get("H3_CONDITIONER", "diffusers-internal-dev/minimax-h3-conditioner")
33
+ # `resident` keeps the 61.73 GiB transformer and the ~20.5 GiB of float32 VAEs on the card at once (82.3 of 95.0 GiB,
34
+ # leaving ~12.7 GiB for activations); `offload` hands placement to `ComponentsManager.enable_auto_cpu_offload`.
35
+ PLACEMENT = os.environ.get("H3_PLACEMENT", "resident").lower()
36
+ # cuDNN's fused attention is 10-20% faster than the SDPA default on this pool and needs nothing installed.
37
+ # flash-attention 3 is sm90-only and this card is sm120 (the `zero-a10g` flavour name is legacy).
38
+ ATTENTION = os.environ.get("H3_ATTENTION", "_native_cudnn").lower()
39
+ GPU_DURATION = int(os.environ.get("H3_GPU_DURATION", "900"))
40
+ GPU_SIZE = os.environ.get("H3_GPU_SIZE", "xlarge")
41
+ ON_SPACES = bool(os.environ.get("SPACE_ID"))
42
+
43
+ CANVASES = {
44
+ "16:9 (768x1344)": (768, 1344),
45
+ "9:16 (1344x768)": (1344, 768),
46
+ "1:1 (768x768)": (768, 768),
47
+ "4:3 (768x1024)": (768, 1024),
48
+ "3:4 (1024x768)": (1024, 768),
49
+ "21:9 (672x1536)": (672, 1536),
50
+ }
51
+ DEFAULT_CANVAS = "16:9 (768x1344)"
52
+ FPS, FRAMES_PER_CHUNK, LATENTS_PER_CHUNK = 24, 17, 5
53
+ MAX_UI_DURATION = 14
54
+
55
+
56
+ def snap_frames(seconds: float) -> int:
57
+ """The frame count MiniMax-H3's video VAE can decode: the next `17 * n + 5` at 24 fps."""
58
+ frames = max(1, round(float(seconds) * FPS))
59
+ while frames % FRAMES_PER_CHUNK != LATENTS_PER_CHUNK:
60
+ frames += 1
61
+ return frames
62
+
63
+
64
+ PIPE = None
65
+ MANAGER = None
66
+ LOAD_ERROR: str | None = None
67
+ LOADED_IN: float | None = None
68
+ CLIENT = None
69
+
70
+
71
+ def status() -> str:
72
+ if LOAD_ERROR:
73
+ return LOAD_ERROR
74
+ if PIPE is None:
75
+ return f"Loading `{MODEL_REPO}` (transformer + VAEs, 77.3 GB). Watch the Space logs."
76
+ return (
77
+ f"Ready · transformer + VAEs **bfloat16, unquantized** · placement `{PLACEMENT}` · attention `{ATTENTION}` · "
78
+ f"loaded in {LOADED_IN:.0f}s · conditioner `{CONDITIONER_SPACE}`"
79
+ )
80
+
81
+
82
+ def load_models() -> str | None:
83
+ """Load the denoising half. At **startup**.
84
+
85
+ `MiniMaxH3GeneratorBlocks` declares `transformer`, `vae`, `audio_vae`, `scheduler`, `audio_scheduler` and
86
+ `video_processor`, so `load_components` fetches exactly those subfolders out of the shared
87
+ `modular_model_index.json` — `text_encoder/` and `transformer_ref/` are never touched.
88
+
89
+ Both autoencoders carry `_keep_in_fp32_modules` over every module, so the `dtype` below is refused for them and
90
+ they load float32 (~20.5 GiB rather than 10.26): a bfloat16 audio VAE decodes the soundtrack ~20 dB too quiet.
91
+ """
92
+ global PIPE, MANAGER, LOAD_ERROR, LOADED_IN
93
+
94
+ if PIPE is not None or LOAD_ERROR is not None:
95
+ return LOAD_ERROR
96
+
97
+ token = os.environ.get("HF_TOKEN")
98
+ if not token:
99
+ LOAD_ERROR = f"**`HF_TOKEN` secret is missing** and `{MODEL_REPO}` is private. Add it and restart."
100
+ return LOAD_ERROR
101
+
102
+ started = time.time()
103
+ try:
104
+ import torch
105
+ from diffusers import ComponentsManager
106
+
107
+ from h3_split_blocks import MiniMaxH3GeneratorBlocks
108
+
109
+ manager = ComponentsManager()
110
+ blocks = MiniMaxH3GeneratorBlocks()
111
+ print(f"[gen] loading {[c.name for c in blocks.expected_components]} from {MODEL_REPO} ...", flush=True)
112
+ pipe = blocks.init_pipeline(MODEL_REPO, components_manager=manager, collection="h3")
113
+ pipe.load_components(dtype=torch.bfloat16, token=token)
114
+ pipe.transformer.set_attention_backend(ATTENTION)
115
+
116
+ if PLACEMENT == "resident":
117
+ # Plain bfloat16 tensors, so ZeroGPU's startup packing handles them — the thing that cannot be moved at
118
+ # startup is a torchao `Float8Tensor`, whose `aten.empty_like(pin_memory=True)` is unimplemented.
119
+ pipe.to("cuda")
120
+ else:
121
+ manager.enable_auto_cpu_offload(device="cuda")
122
+ _arm_decode_hooks(pipe)
123
+
124
+ PIPE, MANAGER = pipe, manager
125
+ LOADED_IN = time.time() - started
126
+ print(f"[gen] ready in {LOADED_IN:.0f}s", flush=True)
127
+ except Exception as error:
128
+ traceback.print_exc()
129
+ LOAD_ERROR = f"**Loading `{MODEL_REPO}` failed** after {time.time() - started:.0f}s: `{type(error).__name__}: {error}`"
130
+ return LOAD_ERROR
131
+
132
+
133
+ def _arm_decode_hooks(pipe):
134
+ """Make the offload hooks fire for the two VAEs.
135
+
136
+ `enable_auto_cpu_offload` installs accelerate hooks, which wrap `forward`. The decode blocks call
137
+ `components.vae.decode(...)` and `components.audio_vae.decode(...)` directly, so the hook never runs and the VAE
138
+ is still on the host when the latents arrive on the card.
139
+ """
140
+ for name in ("vae", "audio_vae"):
141
+ module = getattr(pipe, name)
142
+ inner = module.decode
143
+
144
+ def armed(*args, _module=module, _decode=inner, **kwargs):
145
+ hook = getattr(_module, "_hf_hook", None)
146
+ if hook is not None:
147
+ hook.pre_forward(_module)
148
+ return _decode(*args, **kwargs)
149
+
150
+ module.decode = armed
151
+
152
+
153
+ def conditioner():
154
+ """The other half, over the gradio API. Cached — building a `Client` costs a round trip to the Space config."""
155
+ global CLIENT
156
+ if CLIENT is None:
157
+ from gradio_client import Client
158
+
159
+ CLIENT = Client(CONDITIONER_SPACE, token=os.environ.get("HF_TOKEN"))
160
+ return CLIENT
161
+
162
+
163
+ def encode_remote(prompt, image_path, last_image_path, canvas, num_frames):
164
+ """Ask the conditioner Space for `prompt_embeds` + `text_token_tags`. Off this Space's GPU time entirely."""
165
+ from gradio_client import handle_file
166
+ from safetensors import safe_open
167
+
168
+ path, plan = conditioner().predict(
169
+ prompt=prompt,
170
+ image_path=handle_file(image_path) if image_path else None,
171
+ last_image_path=handle_file(last_image_path) if last_image_path else None,
172
+ canvas=canvas,
173
+ num_frames=num_frames,
174
+ api_name="/encode",
175
+ )
176
+ with safe_open(path, framework="pt") as handle:
177
+ metadata = handle.metadata()
178
+ return handle.get_tensor("prompt_embeds"), handle.get_tensor("text_token_tags"), metadata, plan
179
+
180
+
181
+ @spaces.GPU(duration=GPU_DURATION, size=GPU_SIZE)
182
+ def _generate(prompt_embeds, text_token_tags, image, last_image, height, width, num_frames, steps, seed):
183
+ """The only thing on GPU time: the packed-sequence denoise loop and the two decoders."""
184
+ import torch
185
+
186
+ return PIPE(
187
+ prompt_embeds=prompt_embeds.to("cuda"),
188
+ text_token_tags=text_token_tags,
189
+ image=image,
190
+ last_image=last_image,
191
+ height=height,
192
+ width=width,
193
+ num_frames=num_frames,
194
+ num_inference_steps=int(steps),
195
+ generator=torch.Generator("cpu").manual_seed(int(seed)),
196
+ )
197
+
198
+
199
+ def generate(prompt, image_path, last_image_path, canvas, duration, steps, seed, progress=gr.Progress()):
200
+ if LOAD_ERROR:
201
+ raise gr.Error(LOAD_ERROR)
202
+ if PIPE is None:
203
+ raise gr.Error("The denoiser is still loading.")
204
+ if not prompt or not prompt.strip():
205
+ raise gr.Error("MiniMax-H3 always takes a prompt, keyframes or not.")
206
+
207
+ from PIL import Image
208
+
209
+ from diffusers.utils import encode_video
210
+
211
+ num_frames = snap_frames(duration)
212
+
213
+ progress(0.0, desc=f"Conditioning on {CONDITIONER_SPACE} ...")
214
+ conditioned = time.time()
215
+ prompt_embeds, text_token_tags, metadata, plan = encode_remote(
216
+ prompt, image_path, last_image_path, canvas, num_frames
217
+ )
218
+ condition_seconds = time.time() - conditioned
219
+ height, width, num_frames = (int(metadata[key]) for key in ("height", "width", "num_frames"))
220
+
221
+ progress(0.1, desc=f"Denoising {steps} steps at {width}x{height}, {num_frames} frames ...")
222
+ started = time.time()
223
+ state = _generate(
224
+ prompt_embeds,
225
+ text_token_tags,
226
+ Image.open(image_path) if image_path else None,
227
+ Image.open(last_image_path) if last_image_path else None,
228
+ height,
229
+ width,
230
+ num_frames,
231
+ steps,
232
+ seed,
233
+ )
234
+ generate_seconds = time.time() - started
235
+
236
+ directory = os.path.join(tempfile.gettempdir(), "h3-outputs")
237
+ os.makedirs(directory, exist_ok=True)
238
+ path = os.path.join(directory, f"h3-{int(time.time() * 1000)}.mp4")
239
+ encode_video(
240
+ state.get("videos")[0],
241
+ fps=FPS,
242
+ output_path=path,
243
+ audio=state.get("audio")[0],
244
+ audio_sample_rate=state.get("sampling_rate"),
245
+ )
246
+
247
+ report = (
248
+ f"`{width}x{height}`, {num_frames} frames ({num_frames / FPS:.3f} s), {int(steps)} steps · "
249
+ f"conditioner {condition_seconds:.0f}s ({plan['num_text_tokens']} tokens) · "
250
+ f"denoise + decode {generate_seconds:.0f}s ({generate_seconds / int(steps):.1f} s/step) · seed {int(seed)}"
251
+ )
252
+ print(f"[gen] {report}", flush=True)
253
+ return path, report
254
+
255
+
256
+ load_models()
257
+
258
+ INTRO = """# MiniMax-H3 — unquantized, split across two Spaces
259
+
260
+ Joint video **and** soundtrack out of one denoising pass, at **bfloat16, no quantization anywhere**.
261
+
262
+ MiniMax-H3 is 195.9 GiB in bfloat16 and a ZeroGPU Space is evicted at 150 GB of storage, so the unquantized
263
+ checkpoint does not fit in one Space. It does fit in two: the 62.14 GiB Qwen3-VL conditioner runs in
264
+ [`minimax-h3-conditioner`](https://huggingface.co/spaces/diffusers-internal-dev/minimax-h3-conditioner) and this
265
+ Space holds the 61.73 GiB transformer plus the two autoencoders. Every request calls the conditioner over the gradio
266
+ API and gets back `prompt_embeds` `(1, num_text_tokens, 5120)` and `text_token_tags` `(num_text_tokens,)` — the whole
267
+ wire format of the split.
268
+
269
+ Fixed by the checkpoint: 24 fps, a 768 pixel short edge, 5–15 s, no CFG and no negative prompt.
270
+ """
271
+
272
+
273
+ with gr.Blocks(title="MiniMax-H3 (split, bf16)") as demo:
274
+ gr.Markdown(INTRO)
275
+ banner = gr.Markdown(status())
276
+
277
+ with gr.Row():
278
+ with gr.Column():
279
+ prompt = gr.Textbox(
280
+ label="Prompt",
281
+ lines=3,
282
+ value="A red fox trotting through a snowy pine forest at dawn, snow crunching underfoot",
283
+ )
284
+ with gr.Row():
285
+ image = gr.Image(label="First keyframe (optional)", type="filepath")
286
+ last_image = gr.Image(label="Last keyframe (optional)", type="filepath")
287
+ canvas = gr.Dropdown(label="Canvas", choices=list(CANVASES), value=DEFAULT_CANVAS)
288
+ duration = gr.Slider(label="Duration (s)", minimum=5, maximum=MAX_UI_DURATION, step=1, value=5)
289
+ steps = gr.Slider(label="Steps", minimum=10, maximum=40, step=1, value=30)
290
+ seed = gr.Number(label="Seed", value=42, precision=0)
291
+ run = gr.Button("Generate", variant="primary")
292
+ with gr.Column():
293
+ video = gr.Video(label="Video + soundtrack")
294
+ report = gr.Markdown()
295
+
296
+ run.click(
297
+ generate,
298
+ [prompt, image, last_image, canvas, duration, steps, seed],
299
+ [video, report],
300
+ api_name="generate",
301
+ )
302
+ demo.load(status, None, banner, api_name="status")
303
+
304
+
305
+ if __name__ == "__main__":
306
+ demo.queue(max_size=4).launch(show_error=True)
diffusers/__init__.py ADDED
@@ -0,0 +1,1750 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ __version__ = "0.40.0.dev0"
2
+
3
+ from typing import TYPE_CHECKING
4
+
5
+ from .utils import (
6
+ DIFFUSERS_SLOW_IMPORT,
7
+ OptionalDependencyNotAvailable,
8
+ _LazyModule,
9
+ is_accelerate_available,
10
+ is_auto_round_available,
11
+ is_bitsandbytes_available,
12
+ is_gguf_available,
13
+ is_librosa_available,
14
+ is_note_seq_available,
15
+ is_nvidia_modelopt_available,
16
+ is_onnx_available,
17
+ is_opencv_available,
18
+ is_optimum_quanto_available,
19
+ is_scipy_available,
20
+ is_sdnq_available,
21
+ is_sentencepiece_available,
22
+ is_torch_available,
23
+ is_torchao_available,
24
+ is_torchsde_available,
25
+ is_transformers_available,
26
+ is_transformers_version,
27
+ )
28
+
29
+
30
+ # Lazy Import based on
31
+ # https://github.com/huggingface/transformers/blob/main/src/transformers/__init__.py
32
+
33
+ # When adding a new object to this init, please add it to `_import_structure`. The `_import_structure` is a dictionary submodule to list of object names,
34
+ # and is used to defer the actual importing for when the objects are requested.
35
+ # This way `import diffusers` provides the names in the namespace without actually importing anything (and especially none of the backends).
36
+
37
+ _import_structure = {
38
+ "configuration_utils": ["ConfigMixin"],
39
+ "guiders": [],
40
+ "hooks": [],
41
+ "loaders": ["FromOriginalModelMixin"],
42
+ "models": [],
43
+ "modular_pipelines": [],
44
+ "pipelines": [],
45
+ "quantizers.pipe_quant_config": ["PipelineQuantizationConfig"],
46
+ "quantizers.quantization_config": [],
47
+ "schedulers": [],
48
+ "utils": [
49
+ "OptionalDependencyNotAvailable",
50
+ "is_inflect_available",
51
+ "is_invisible_watermark_available",
52
+ "is_librosa_available",
53
+ "is_note_seq_available",
54
+ "is_onnx_available",
55
+ "is_scipy_available",
56
+ "is_torch_available",
57
+ "is_torchsde_available",
58
+ "is_transformers_available",
59
+ "is_transformers_version",
60
+ "is_unidecode_available",
61
+ "logging",
62
+ ],
63
+ }
64
+
65
+ try:
66
+ if not is_torch_available() and not is_accelerate_available() and not is_bitsandbytes_available():
67
+ raise OptionalDependencyNotAvailable()
68
+ except OptionalDependencyNotAvailable:
69
+ from .utils import dummy_bitsandbytes_objects
70
+
71
+ _import_structure["utils.dummy_bitsandbytes_objects"] = [
72
+ name for name in dir(dummy_bitsandbytes_objects) if not name.startswith("_")
73
+ ]
74
+ else:
75
+ _import_structure["quantizers.quantization_config"].append("BitsAndBytesConfig")
76
+
77
+ try:
78
+ if not is_torch_available() and not is_accelerate_available() and not is_gguf_available():
79
+ raise OptionalDependencyNotAvailable()
80
+ except OptionalDependencyNotAvailable:
81
+ from .utils import dummy_gguf_objects
82
+
83
+ _import_structure["utils.dummy_gguf_objects"] = [
84
+ name for name in dir(dummy_gguf_objects) if not name.startswith("_")
85
+ ]
86
+ else:
87
+ _import_structure["quantizers.quantization_config"].append("GGUFQuantizationConfig")
88
+
89
+ try:
90
+ if not is_torch_available() and not is_accelerate_available() and not is_torchao_available():
91
+ raise OptionalDependencyNotAvailable()
92
+ except OptionalDependencyNotAvailable:
93
+ from .utils import dummy_torchao_objects
94
+
95
+ _import_structure["utils.dummy_torchao_objects"] = [
96
+ name for name in dir(dummy_torchao_objects) if not name.startswith("_")
97
+ ]
98
+ else:
99
+ _import_structure["quantizers.quantization_config"].append("TorchAoConfig")
100
+
101
+ try:
102
+ if not is_torch_available() and not is_accelerate_available() and not is_optimum_quanto_available():
103
+ raise OptionalDependencyNotAvailable()
104
+ except OptionalDependencyNotAvailable:
105
+ from .utils import dummy_optimum_quanto_objects
106
+
107
+ _import_structure["utils.dummy_optimum_quanto_objects"] = [
108
+ name for name in dir(dummy_optimum_quanto_objects) if not name.startswith("_")
109
+ ]
110
+ else:
111
+ _import_structure["quantizers.quantization_config"].append("QuantoConfig")
112
+
113
+ try:
114
+ if not is_torch_available() and not is_accelerate_available() and not is_nvidia_modelopt_available():
115
+ raise OptionalDependencyNotAvailable()
116
+ except OptionalDependencyNotAvailable:
117
+ from .utils import dummy_nvidia_modelopt_objects
118
+
119
+ _import_structure["utils.dummy_nvidia_modelopt_objects"] = [
120
+ name for name in dir(dummy_nvidia_modelopt_objects) if not name.startswith("_")
121
+ ]
122
+ else:
123
+ _import_structure["quantizers.quantization_config"].append("NVIDIAModelOptConfig")
124
+
125
+ try:
126
+ if not is_torch_available():
127
+ raise OptionalDependencyNotAvailable()
128
+ except OptionalDependencyNotAvailable:
129
+ from .utils import dummy_nunchaku_lite_objects
130
+
131
+ _import_structure["utils.dummy_nunchaku_lite_objects"] = [
132
+ name for name in dir(dummy_nunchaku_lite_objects) if not name.startswith("_")
133
+ ]
134
+ else:
135
+ _import_structure["quantizers.quantization_config"].append("NunchakuLiteQuantizationConfig")
136
+
137
+ try:
138
+ if not is_auto_round_available():
139
+ raise OptionalDependencyNotAvailable()
140
+ except OptionalDependencyNotAvailable:
141
+ from .utils import dummy_auto_round_objects
142
+
143
+ _import_structure["utils.dummy_auto_round_objects"] = [
144
+ name for name in dir(dummy_auto_round_objects) if not name.startswith("_")
145
+ ]
146
+ else:
147
+ _import_structure["quantizers.quantization_config"].append("AutoRoundConfig")
148
+
149
+ try:
150
+ if not is_torch_available() and not is_accelerate_available() and not is_sdnq_available():
151
+ raise OptionalDependencyNotAvailable()
152
+ except OptionalDependencyNotAvailable:
153
+ from .utils import dummy_sdnq_objects
154
+
155
+ _import_structure["utils.dummy_sdnq_objects"] = [
156
+ name for name in dir(dummy_sdnq_objects) if not name.startswith("_")
157
+ ]
158
+ else:
159
+ _import_structure["quantizers.quantization_config"].append("SDNQConfig")
160
+
161
+ try:
162
+ if not is_onnx_available():
163
+ raise OptionalDependencyNotAvailable()
164
+ except OptionalDependencyNotAvailable:
165
+ from .utils import dummy_onnx_objects # noqa F403
166
+
167
+ _import_structure["utils.dummy_onnx_objects"] = [
168
+ name for name in dir(dummy_onnx_objects) if not name.startswith("_")
169
+ ]
170
+
171
+ else:
172
+ _import_structure["pipelines"].extend(["OnnxRuntimeModel"])
173
+
174
+ try:
175
+ if not is_torch_available():
176
+ raise OptionalDependencyNotAvailable()
177
+ except OptionalDependencyNotAvailable:
178
+ from .utils import dummy_pt_objects # noqa F403
179
+
180
+ _import_structure["utils.dummy_pt_objects"] = [name for name in dir(dummy_pt_objects) if not name.startswith("_")]
181
+
182
+ else:
183
+ _import_structure["guiders"].extend(
184
+ [
185
+ "AdaptiveProjectedGuidance",
186
+ "AdaptiveProjectedMixGuidance",
187
+ "AutoGuidance",
188
+ "BaseGuidance",
189
+ "ClassifierFreeGuidance",
190
+ "ClassifierFreeZeroStarGuidance",
191
+ "FrequencyDecoupledGuidance",
192
+ "PerturbedAttentionGuidance",
193
+ "SkipLayerGuidance",
194
+ "SmoothedEnergyGuidance",
195
+ "TangentialClassifierFreeGuidance",
196
+ ]
197
+ )
198
+ _import_structure["hooks"].extend(
199
+ [
200
+ "FasterCacheConfig",
201
+ "FirstBlockCacheConfig",
202
+ "HookRegistry",
203
+ "LayerSkipConfig",
204
+ "MagCacheConfig",
205
+ "PyramidAttentionBroadcastConfig",
206
+ "SmoothedEnergyGuidanceConfig",
207
+ "TaylorSeerCacheConfig",
208
+ "TextKVCacheConfig",
209
+ "apply_faster_cache",
210
+ "apply_first_block_cache",
211
+ "apply_layer_skip",
212
+ "apply_mag_cache",
213
+ "apply_pyramid_attention_broadcast",
214
+ "apply_taylorseer_cache",
215
+ "apply_text_kv_cache",
216
+ ]
217
+ )
218
+ _import_structure["image_processor"] = [
219
+ "InpaintProcessor",
220
+ "IPAdapterMaskProcessor",
221
+ "PixArtImageProcessor",
222
+ "VaeImageProcessor",
223
+ "VaeImageProcessorLDM3D",
224
+ ]
225
+ _import_structure["models"].extend(
226
+ [
227
+ "AceStepTransformer1DModel",
228
+ "AllegroTransformer3DModel",
229
+ "AnimaTextConditioner",
230
+ "AnyFlowFARTransformer3DModel",
231
+ "AnyFlowTransformer3DModel",
232
+ "AsymmetricAutoencoderKL",
233
+ "AttentionBackendName",
234
+ "AuraFlowTransformer2DModel",
235
+ "AutoencoderDC",
236
+ "AutoencoderKL",
237
+ "AutoencoderKLAllegro",
238
+ "AutoencoderKLCogVideoX",
239
+ "AutoencoderKLCosmos",
240
+ "AutoencoderKLFlux2",
241
+ "AutoencoderKLHunyuanImage",
242
+ "AutoencoderKLHunyuanImageRefiner",
243
+ "AutoencoderKLHunyuanVideo",
244
+ "AutoencoderKLHunyuanVideo15",
245
+ "AutoencoderKLKVAE",
246
+ "AutoencoderKLKVAEVideo",
247
+ "AutoencoderKLLTX2Audio",
248
+ "AutoencoderKLLTX2Video",
249
+ "AutoencoderKLLTXVideo",
250
+ "AutoencoderKLMagvit",
251
+ "AutoencoderKLMiniMaxH3",
252
+ "AutoencoderKLMiniMaxH3Audio",
253
+ "AutoencoderKLMochi",
254
+ "AutoencoderKLQwenImage",
255
+ "AutoencoderKLTemporalDecoder",
256
+ "AutoencoderKLWan",
257
+ "AutoencoderOobleck",
258
+ "AutoencoderRAE",
259
+ "AutoencoderTiny",
260
+ "AutoencoderVidTok",
261
+ "AutoModel",
262
+ "BriaFiboTransformer2DModel",
263
+ "BriaTransformer2DModel",
264
+ "CacheMixin",
265
+ "ChromaTransformer2DModel",
266
+ "ChronoEditTransformer3DModel",
267
+ "CogVideoXTransformer3DModel",
268
+ "CogView3PlusTransformer2DModel",
269
+ "CogView4Transformer2DModel",
270
+ "ConsisIDTransformer3DModel",
271
+ "ConsistencyDecoderVAE",
272
+ "ContextParallelConfig",
273
+ "ControlNetModel",
274
+ "ControlNetUnionModel",
275
+ "ControlNetXSAdapter",
276
+ "Cosmos3AVAEAudioTokenizer",
277
+ "Cosmos3OmniTransformer",
278
+ "CosmosControlNetModel",
279
+ "CosmosTransformer3DModel",
280
+ "DiTTransformer2DModel",
281
+ "DreamLiteTransformer2DModel",
282
+ "DreamLiteUNetModel",
283
+ "EasyAnimateTransformer3DModel",
284
+ "ErnieImageTransformer2DModel",
285
+ "Flux2Transformer2DModel",
286
+ "FluxControlNetModel",
287
+ "FluxMultiControlNetModel",
288
+ "FluxTransformer2DModel",
289
+ "GlmImageTransformer2DModel",
290
+ "HeliosTransformer3DModel",
291
+ "HiDreamImageTransformer2DModel",
292
+ "HunyuanDiT2DControlNetModel",
293
+ "HunyuanDiT2DModel",
294
+ "HunyuanDiT2DMultiControlNetModel",
295
+ "HunyuanImageTransformer2DModel",
296
+ "HunyuanVideo15Transformer3DModel",
297
+ "HunyuanVideoFramepackTransformer3DModel",
298
+ "HunyuanVideoTransformer3DModel",
299
+ "I2VGenXLUNet",
300
+ "Ideogram4Transformer2DModel",
301
+ "JoyImageEditPlusTransformer3DModel",
302
+ "JoyImageEditTransformer3DModel",
303
+ "Kandinsky3UNet",
304
+ "Kandinsky5Transformer3DModel",
305
+ "Krea2Transformer2DModel",
306
+ "LatteTransformer3DModel",
307
+ "LongCatAudioDiTTransformer",
308
+ "LongCatAudioDiTVae",
309
+ "LongCatImageTransformer2DModel",
310
+ "LTX2VideoTransformer3DModel",
311
+ "LTXVideoTransformer3DModel",
312
+ "Lumina2Transformer2DModel",
313
+ "LuminaNextDiT2DModel",
314
+ "MiniMaxH3Transformer3DModel",
315
+ "MochiTransformer3DModel",
316
+ "ModelMixin",
317
+ "MotifVideoTransformer3DModel",
318
+ "MotionAdapter",
319
+ "MultiAdapter",
320
+ "MultiControlNetModel",
321
+ "NucleusMoEImageTransformer2DModel",
322
+ "OmniGenTransformer2DModel",
323
+ "OvisImageTransformer2DModel",
324
+ "ParallelConfig",
325
+ "PixArtTransformer2DModel",
326
+ "PriorTransformer",
327
+ "PRXTransformer2DModel",
328
+ "QwenImageControlNetModel",
329
+ "QwenImageMultiControlNetModel",
330
+ "QwenImageTransformer2DModel",
331
+ "SanaControlNetModel",
332
+ "SanaTransformer2DModel",
333
+ "SanaVideoTransformer3DModel",
334
+ "SD3ControlNetModel",
335
+ "SD3MultiControlNetModel",
336
+ "SD3Transformer2DModel",
337
+ "SkyReelsV2Transformer3DModel",
338
+ "SparseControlNetModel",
339
+ "StableAudioDiTModel",
340
+ "StableCascadeUNet",
341
+ "T2IAdapter",
342
+ "T5FilmDecoder",
343
+ "Transformer2DModel",
344
+ "TransformerTemporalModel",
345
+ "UNet1DModel",
346
+ "UNet2DConditionModel",
347
+ "UNet2DModel",
348
+ "UNet3DConditionModel",
349
+ "UNetControlNetXSModel",
350
+ "UNetMotionModel",
351
+ "UNetSpatioTemporalConditionModel",
352
+ "UVit2DModel",
353
+ "VQModel",
354
+ "WanAnimateTransformer3DModel",
355
+ "WanTransformer3DModel",
356
+ "WanVACETransformer3DModel",
357
+ "ZImageControlNetModel",
358
+ "ZImageTransformer2DModel",
359
+ "attention_backend",
360
+ ]
361
+ )
362
+ _import_structure["modular_pipelines"].extend(
363
+ [
364
+ "AutoPipelineBlocks",
365
+ "ComponentsManager",
366
+ "ComponentSpec",
367
+ "ConditionalPipelineBlocks",
368
+ "ConfigSpec",
369
+ "InputParam",
370
+ "LoopSequentialPipelineBlocks",
371
+ "ModularPipeline",
372
+ "ModularPipelineBlocks",
373
+ "OutputParam",
374
+ "SequentialPipelineBlocks",
375
+ ]
376
+ )
377
+ _import_structure["optimization"] = [
378
+ "get_constant_schedule",
379
+ "get_constant_schedule_with_warmup",
380
+ "get_cosine_schedule_with_warmup",
381
+ "get_cosine_with_hard_restarts_schedule_with_warmup",
382
+ "get_linear_schedule_with_warmup",
383
+ "get_polynomial_decay_schedule_with_warmup",
384
+ "get_scheduler",
385
+ ]
386
+ _import_structure["pipelines"].extend(
387
+ [
388
+ "AudioPipelineOutput",
389
+ "AutoPipelineForImage2Image",
390
+ "AutoPipelineForInpainting",
391
+ "AutoPipelineForText2Audio",
392
+ "AutoPipelineForText2Image",
393
+ "ConsistencyModelPipeline",
394
+ "DanceDiffusionPipeline",
395
+ "DDIMPipeline",
396
+ "DDPMPipeline",
397
+ "DiffusionPipeline",
398
+ "DiTPipeline",
399
+ "ImagePipelineOutput",
400
+ "KarrasVePipeline",
401
+ "LDMPipeline",
402
+ "LDMSuperResolutionPipeline",
403
+ "PNDMPipeline",
404
+ "RePaintPipeline",
405
+ "ScoreSdeVePipeline",
406
+ "StableDiffusionMixin",
407
+ ]
408
+ )
409
+ _import_structure["quantizers"] = ["DiffusersQuantizer"]
410
+ _import_structure["schedulers"].extend(
411
+ [
412
+ "AmusedScheduler",
413
+ "BlockRefinementScheduler",
414
+ "BlockRefinementSchedulerOutput",
415
+ "CMStochasticIterativeScheduler",
416
+ "CogVideoXDDIMScheduler",
417
+ "CogVideoXDPMScheduler",
418
+ "DDIMInverseScheduler",
419
+ "DDIMParallelScheduler",
420
+ "DDIMScheduler",
421
+ "DDPMParallelScheduler",
422
+ "DDPMScheduler",
423
+ "DDPMWuerstchenScheduler",
424
+ "DEISMultistepScheduler",
425
+ "DiscreteDDIMScheduler",
426
+ "DiscreteDDIMSchedulerOutput",
427
+ "DPMSolverMultistepInverseScheduler",
428
+ "DPMSolverMultistepScheduler",
429
+ "DPMSolverSinglestepScheduler",
430
+ "EDMDPMSolverMultistepScheduler",
431
+ "EDMEulerScheduler",
432
+ "EntropyBoundScheduler",
433
+ "EntropyBoundSchedulerOutput",
434
+ "EulerAncestralDiscreteScheduler",
435
+ "EulerDiscreteScheduler",
436
+ "FlowMapEulerDiscreteScheduler",
437
+ "FlowMatchEulerDiscreteScheduler",
438
+ "FlowMatchHeunDiscreteScheduler",
439
+ "FlowMatchLCMScheduler",
440
+ "HeliosDMDScheduler",
441
+ "HeliosScheduler",
442
+ "HeunDiscreteScheduler",
443
+ "IPNDMScheduler",
444
+ "KarrasVeScheduler",
445
+ "KDPM2AncestralDiscreteScheduler",
446
+ "KDPM2DiscreteScheduler",
447
+ "LCMScheduler",
448
+ "LTXEulerAncestralRFScheduler",
449
+ "MiniMaxH3Scheduler",
450
+ "PNDMScheduler",
451
+ "RePaintScheduler",
452
+ "SASolverScheduler",
453
+ "SchedulerMixin",
454
+ "SCMScheduler",
455
+ "ScoreSdeVeScheduler",
456
+ "TCDScheduler",
457
+ "UnCLIPScheduler",
458
+ "UniPCMultistepScheduler",
459
+ "VQDiffusionScheduler",
460
+ ]
461
+ )
462
+ _import_structure["training_utils"] = ["EMAModel"]
463
+ _import_structure["video_processor"] = ["VideoProcessor"]
464
+
465
+ try:
466
+ if not (is_torch_available() and is_scipy_available()):
467
+ raise OptionalDependencyNotAvailable()
468
+ except OptionalDependencyNotAvailable:
469
+ from .utils import dummy_torch_and_scipy_objects # noqa F403
470
+
471
+ _import_structure["utils.dummy_torch_and_scipy_objects"] = [
472
+ name for name in dir(dummy_torch_and_scipy_objects) if not name.startswith("_")
473
+ ]
474
+
475
+ else:
476
+ _import_structure["schedulers"].extend(["LMSDiscreteScheduler"])
477
+
478
+ try:
479
+ if not (is_torch_available() and is_torchsde_available()):
480
+ raise OptionalDependencyNotAvailable()
481
+ except OptionalDependencyNotAvailable:
482
+ from .utils import dummy_torch_and_torchsde_objects # noqa F403
483
+
484
+ _import_structure["utils.dummy_torch_and_torchsde_objects"] = [
485
+ name for name in dir(dummy_torch_and_torchsde_objects) if not name.startswith("_")
486
+ ]
487
+
488
+ else:
489
+ _import_structure["schedulers"].extend(["CosineDPMSolverMultistepScheduler", "DPMSolverSDEScheduler"])
490
+
491
+ try:
492
+ if not (is_torch_available() and is_transformers_available()):
493
+ raise OptionalDependencyNotAvailable()
494
+ except OptionalDependencyNotAvailable:
495
+ from .utils import dummy_torch_and_transformers_objects # noqa F403
496
+
497
+ _import_structure["utils.dummy_torch_and_transformers_objects"] = [
498
+ name for name in dir(dummy_torch_and_transformers_objects) if not name.startswith("_")
499
+ ]
500
+
501
+ else:
502
+ _import_structure["modular_pipelines"].extend(
503
+ [
504
+ "AnimaAutoBlocks",
505
+ "AnimaModularPipeline",
506
+ "Cosmos3DistilledBlocks",
507
+ "Cosmos3DistilledModularPipeline",
508
+ "Cosmos3OmniBlocks",
509
+ "Cosmos3OmniModularPipeline",
510
+ "ErnieImageAutoBlocks",
511
+ "ErnieImageModularPipeline",
512
+ "Flux2AutoBlocks",
513
+ "Flux2KleinAutoBlocks",
514
+ "Flux2KleinBaseAutoBlocks",
515
+ "Flux2KleinBaseModularPipeline",
516
+ "Flux2KleinModularPipeline",
517
+ "Flux2ModularPipeline",
518
+ "FluxAutoBlocks",
519
+ "FluxKontextAutoBlocks",
520
+ "FluxKontextModularPipeline",
521
+ "FluxModularPipeline",
522
+ "HeliosAutoBlocks",
523
+ "HeliosModularPipeline",
524
+ "HeliosPyramidAutoBlocks",
525
+ "HeliosPyramidDistilledAutoBlocks",
526
+ "HeliosPyramidDistilledModularPipeline",
527
+ "HeliosPyramidModularPipeline",
528
+ "HunyuanVideo15AutoBlocks",
529
+ "HunyuanVideo15ModularPipeline",
530
+ "Ideogram4AutoBlocks",
531
+ "Ideogram4ModularPipeline",
532
+ "Krea2AutoBlocks",
533
+ "Krea2ModularPipeline",
534
+ "Krea2TurboAutoBlocks",
535
+ "Krea2TurboModularPipeline",
536
+ "LTXAutoBlocks",
537
+ "LTXModularPipeline",
538
+ "MiniMaxH3Blocks",
539
+ "MiniMaxH3ModularPipeline",
540
+ "MiniMaxH3Ref2VABlocks",
541
+ "MiniMaxH3Ref2VAModularPipeline",
542
+ "QwenImageAutoBlocks",
543
+ "QwenImageEditAutoBlocks",
544
+ "QwenImageEditModularPipeline",
545
+ "QwenImageEditPlusAutoBlocks",
546
+ "QwenImageEditPlusModularPipeline",
547
+ "QwenImageLayeredAutoBlocks",
548
+ "QwenImageLayeredModularPipeline",
549
+ "QwenImageModularPipeline",
550
+ "StableDiffusion3AutoBlocks",
551
+ "StableDiffusion3ModularPipeline",
552
+ "StableDiffusionXLAutoBlocks",
553
+ "StableDiffusionXLModularPipeline",
554
+ "Wan22Blocks",
555
+ "Wan22Image2VideoBlocks",
556
+ "Wan22Image2VideoModularPipeline",
557
+ "Wan22ModularPipeline",
558
+ "WanBlocks",
559
+ "WanImage2VideoAutoBlocks",
560
+ "WanImage2VideoModularPipeline",
561
+ "WanModularPipeline",
562
+ "ZImageAutoBlocks",
563
+ "ZImageModularPipeline",
564
+ ]
565
+ )
566
+ _import_structure["pipelines"].extend(
567
+ [
568
+ "AceStepAudioTokenDetokenizer",
569
+ "AceStepAudioTokenizer",
570
+ "AceStepConditionEncoder",
571
+ "AceStepPipeline",
572
+ "AllegroPipeline",
573
+ "AltDiffusionImg2ImgPipeline",
574
+ "AltDiffusionPipeline",
575
+ "AmusedImg2ImgPipeline",
576
+ "AmusedInpaintPipeline",
577
+ "AmusedPipeline",
578
+ "AnimateDiffControlNetPipeline",
579
+ "AnimateDiffPAGPipeline",
580
+ "AnimateDiffPipeline",
581
+ "AnimateDiffSDXLPipeline",
582
+ "AnimateDiffSparseControlNetPipeline",
583
+ "AnimateDiffVideoToVideoControlNetPipeline",
584
+ "AnimateDiffVideoToVideoPipeline",
585
+ "AnyFlowFARPipeline",
586
+ "AnyFlowPipeline",
587
+ "AudioLDM2Pipeline",
588
+ "AudioLDM2ProjectionModel",
589
+ "AudioLDM2UNet2DConditionModel",
590
+ "AudioLDMPipeline",
591
+ "AuraFlowPipeline",
592
+ "BlipDiffusionControlNetPipeline",
593
+ "BlipDiffusionPipeline",
594
+ "BriaFiboEditPipeline",
595
+ "BriaFiboPipeline",
596
+ "BriaPipeline",
597
+ "ChromaImg2ImgPipeline",
598
+ "ChromaInpaintPipeline",
599
+ "ChromaPipeline",
600
+ "ChronoEditPipeline",
601
+ "CLIPImageProjection",
602
+ "CogVideoXFunControlPipeline",
603
+ "CogVideoXImageToVideoPipeline",
604
+ "CogVideoXPipeline",
605
+ "CogVideoXVideoToVideoPipeline",
606
+ "CogView3PlusPipeline",
607
+ "CogView4ControlPipeline",
608
+ "CogView4Pipeline",
609
+ "ConsisIDPipeline",
610
+ "Cosmos2_5_PredictBasePipeline",
611
+ "Cosmos2_5_TransferPipeline",
612
+ "Cosmos2TextToImagePipeline",
613
+ "Cosmos2VideoToWorldPipeline",
614
+ "Cosmos3OmniPipeline",
615
+ "CosmosActionCondition",
616
+ "CosmosTextToWorldPipeline",
617
+ "CosmosVideoToWorldPipeline",
618
+ "CycleDiffusionPipeline",
619
+ "DiffusionGemmaPipeline",
620
+ "DiffusionGemmaPipelineOutput",
621
+ "DreamLiteMobilePipeline",
622
+ "DreamLitePipeline",
623
+ "DreamLitePipelineOutput",
624
+ "EasyAnimateControlPipeline",
625
+ "EasyAnimateInpaintPipeline",
626
+ "EasyAnimatePipeline",
627
+ "ErnieImagePipeline",
628
+ "Flux2KleinInpaintPipeline",
629
+ "Flux2KleinKVPipeline",
630
+ "Flux2KleinPipeline",
631
+ "Flux2Pipeline",
632
+ "FluxControlImg2ImgPipeline",
633
+ "FluxControlInpaintPipeline",
634
+ "FluxControlNetImg2ImgPipeline",
635
+ "FluxControlNetInpaintPipeline",
636
+ "FluxControlNetPipeline",
637
+ "FluxControlPipeline",
638
+ "FluxFillPipeline",
639
+ "FluxImg2ImgPipeline",
640
+ "FluxInpaintPipeline",
641
+ "FluxKontextInpaintPipeline",
642
+ "FluxKontextPipeline",
643
+ "FluxPipeline",
644
+ "FluxPriorReduxPipeline",
645
+ "GlmImagePipeline",
646
+ "HeliosPipeline",
647
+ "HeliosPyramidPipeline",
648
+ "HiDreamImagePipeline",
649
+ "HunyuanDiTControlNetPipeline",
650
+ "HunyuanDiTPAGPipeline",
651
+ "HunyuanDiTPipeline",
652
+ "HunyuanImagePipeline",
653
+ "HunyuanImageRefinerPipeline",
654
+ "HunyuanSkyreelsImageToVideoPipeline",
655
+ "HunyuanVideo15ImageToVideoPipeline",
656
+ "HunyuanVideo15Pipeline",
657
+ "HunyuanVideoFramepackPipeline",
658
+ "HunyuanVideoImageToVideoPipeline",
659
+ "HunyuanVideoPipeline",
660
+ "I2VGenXLPipeline",
661
+ "Ideogram4Pipeline",
662
+ "Ideogram4PromptEnhancerHead",
663
+ "IFImg2ImgPipeline",
664
+ "IFImg2ImgSuperResolutionPipeline",
665
+ "IFInpaintingPipeline",
666
+ "IFInpaintingSuperResolutionPipeline",
667
+ "IFPipeline",
668
+ "IFSuperResolutionPipeline",
669
+ "ImageTextPipelineOutput",
670
+ "JoyImageEditPipeline",
671
+ "JoyImageEditPipelineOutput",
672
+ "JoyImageEditPlusPipeline",
673
+ "JoyImageEditPlusPipelineOutput",
674
+ "Kandinsky3Img2ImgPipeline",
675
+ "Kandinsky3Pipeline",
676
+ "Kandinsky5I2IPipeline",
677
+ "Kandinsky5I2VPipeline",
678
+ "Kandinsky5T2IPipeline",
679
+ "Kandinsky5T2VPipeline",
680
+ "KandinskyCombinedPipeline",
681
+ "KandinskyImg2ImgCombinedPipeline",
682
+ "KandinskyImg2ImgPipeline",
683
+ "KandinskyInpaintCombinedPipeline",
684
+ "KandinskyInpaintPipeline",
685
+ "KandinskyPipeline",
686
+ "KandinskyPriorPipeline",
687
+ "KandinskyV22CombinedPipeline",
688
+ "KandinskyV22ControlnetImg2ImgPipeline",
689
+ "KandinskyV22ControlnetPipeline",
690
+ "KandinskyV22Img2ImgCombinedPipeline",
691
+ "KandinskyV22Img2ImgPipeline",
692
+ "KandinskyV22InpaintCombinedPipeline",
693
+ "KandinskyV22InpaintPipeline",
694
+ "KandinskyV22Pipeline",
695
+ "KandinskyV22PriorEmb2EmbPipeline",
696
+ "KandinskyV22PriorPipeline",
697
+ "Krea2Pipeline",
698
+ "LatentConsistencyModelImg2ImgPipeline",
699
+ "LatentConsistencyModelPipeline",
700
+ "LattePipeline",
701
+ "LDMTextToImagePipeline",
702
+ "LEditsPPPipelineStableDiffusion",
703
+ "LEditsPPPipelineStableDiffusionXL",
704
+ "LLaDA2Pipeline",
705
+ "LLaDA2PipelineOutput",
706
+ "LongCatAudioDiTPipeline",
707
+ "LongCatImageEditPipeline",
708
+ "LongCatImagePipeline",
709
+ "LTX2ConditionPipeline",
710
+ "LTX2HDRPipeline",
711
+ "LTX2ImageToVideoPipeline",
712
+ "LTX2InContextPipeline",
713
+ "LTX2LatentUpsamplePipeline",
714
+ "LTX2Pipeline",
715
+ "LTXConditionPipeline",
716
+ "LTXI2VLongMultiPromptPipeline",
717
+ "LTXImageToVideoPipeline",
718
+ "LTXLatentUpsamplePipeline",
719
+ "LTXPipeline",
720
+ "LucyEditPipeline",
721
+ "Lumina2Pipeline",
722
+ "Lumina2Text2ImgPipeline",
723
+ "LuminaPipeline",
724
+ "LuminaText2ImgPipeline",
725
+ "MarigoldDepthPipeline",
726
+ "MarigoldIntrinsicsPipeline",
727
+ "MarigoldNormalsPipeline",
728
+ "MochiPipeline",
729
+ "MotifVideoImage2VideoPipeline",
730
+ "MotifVideoPipeline",
731
+ "MotifVideoPipelineOutput",
732
+ "MusicLDMPipeline",
733
+ "NucleusMoEImagePipeline",
734
+ "OmniGenPipeline",
735
+ "OvisImagePipeline",
736
+ "PaintByExamplePipeline",
737
+ "PIAPipeline",
738
+ "PixArtAlphaPipeline",
739
+ "PixArtSigmaPAGPipeline",
740
+ "PixArtSigmaPipeline",
741
+ "PRXPipeline",
742
+ "PRXPixelPipeline",
743
+ "QwenImageControlNetInpaintPipeline",
744
+ "QwenImageControlNetPipeline",
745
+ "QwenImageEditInpaintPipeline",
746
+ "QwenImageEditPipeline",
747
+ "QwenImageEditPlusPipeline",
748
+ "QwenImageImg2ImgPipeline",
749
+ "QwenImageInpaintPipeline",
750
+ "QwenImageLayeredPipeline",
751
+ "QwenImagePipeline",
752
+ "ReduxImageEncoder",
753
+ "SanaControlNetPipeline",
754
+ "SanaImageToVideoPipeline",
755
+ "SanaPAGPipeline",
756
+ "SanaPipeline",
757
+ "SanaSprintImg2ImgPipeline",
758
+ "SanaSprintPipeline",
759
+ "SanaVideoPipeline",
760
+ "SanaVideoPipeline",
761
+ "SemanticStableDiffusionPipeline",
762
+ "ShapEImg2ImgPipeline",
763
+ "ShapEPipeline",
764
+ "SkyReelsV2DiffusionForcingImageToVideoPipeline",
765
+ "SkyReelsV2DiffusionForcingPipeline",
766
+ "SkyReelsV2DiffusionForcingVideoToVideoPipeline",
767
+ "SkyReelsV2ImageToVideoPipeline",
768
+ "SkyReelsV2Pipeline",
769
+ "StableAudioPipeline",
770
+ "StableAudioProjectionModel",
771
+ "StableCascadeCombinedPipeline",
772
+ "StableCascadeDecoderPipeline",
773
+ "StableCascadePriorPipeline",
774
+ "StableDiffusion3ControlNetInpaintingPipeline",
775
+ "StableDiffusion3ControlNetPipeline",
776
+ "StableDiffusion3Img2ImgPipeline",
777
+ "StableDiffusion3InpaintPipeline",
778
+ "StableDiffusion3PAGImg2ImgPipeline",
779
+ "StableDiffusion3PAGImg2ImgPipeline",
780
+ "StableDiffusion3PAGPipeline",
781
+ "StableDiffusion3Pipeline",
782
+ "StableDiffusionAdapterPipeline",
783
+ "StableDiffusionAttendAndExcitePipeline",
784
+ "StableDiffusionControlNetImg2ImgPipeline",
785
+ "StableDiffusionControlNetInpaintPipeline",
786
+ "StableDiffusionControlNetPAGInpaintPipeline",
787
+ "StableDiffusionControlNetPAGPipeline",
788
+ "StableDiffusionControlNetPipeline",
789
+ "StableDiffusionControlNetXSPipeline",
790
+ "StableDiffusionDepth2ImgPipeline",
791
+ "StableDiffusionDiffEditPipeline",
792
+ "StableDiffusionGLIGENPipeline",
793
+ "StableDiffusionGLIGENTextImagePipeline",
794
+ "StableDiffusionImageVariationPipeline",
795
+ "StableDiffusionImg2ImgPipeline",
796
+ "StableDiffusionInpaintPipeline",
797
+ "StableDiffusionInpaintPipelineLegacy",
798
+ "StableDiffusionInstructPix2PixPipeline",
799
+ "StableDiffusionLatentUpscalePipeline",
800
+ "StableDiffusionLDM3DPipeline",
801
+ "StableDiffusionModelEditingPipeline",
802
+ "StableDiffusionPAGImg2ImgPipeline",
803
+ "StableDiffusionPAGInpaintPipeline",
804
+ "StableDiffusionPAGPipeline",
805
+ "StableDiffusionPanoramaPipeline",
806
+ "StableDiffusionParadigmsPipeline",
807
+ "StableDiffusionPipeline",
808
+ "StableDiffusionPipelineSafe",
809
+ "StableDiffusionPix2PixZeroPipeline",
810
+ "StableDiffusionSAGPipeline",
811
+ "StableDiffusionUpscalePipeline",
812
+ "StableDiffusionXLAdapterPipeline",
813
+ "StableDiffusionXLControlNetImg2ImgPipeline",
814
+ "StableDiffusionXLControlNetInpaintPipeline",
815
+ "StableDiffusionXLControlNetPAGImg2ImgPipeline",
816
+ "StableDiffusionXLControlNetPAGPipeline",
817
+ "StableDiffusionXLControlNetPipeline",
818
+ "StableDiffusionXLControlNetUnionImg2ImgPipeline",
819
+ "StableDiffusionXLControlNetUnionInpaintPipeline",
820
+ "StableDiffusionXLControlNetUnionPipeline",
821
+ "StableDiffusionXLControlNetXSPipeline",
822
+ "StableDiffusionXLImg2ImgPipeline",
823
+ "StableDiffusionXLInpaintPipeline",
824
+ "StableDiffusionXLInstructPix2PixPipeline",
825
+ "StableDiffusionXLPAGImg2ImgPipeline",
826
+ "StableDiffusionXLPAGInpaintPipeline",
827
+ "StableDiffusionXLPAGPipeline",
828
+ "StableDiffusionXLPipeline",
829
+ "StableUnCLIPImg2ImgPipeline",
830
+ "StableUnCLIPPipeline",
831
+ "StableVideoDiffusionPipeline",
832
+ "TextToVideoSDPipeline",
833
+ "TextToVideoZeroPipeline",
834
+ "TextToVideoZeroSDXLPipeline",
835
+ "UnCLIPImageVariationPipeline",
836
+ "UnCLIPPipeline",
837
+ "UniDiffuserModel",
838
+ "UniDiffuserPipeline",
839
+ "UniDiffuserTextDecoder",
840
+ "VersatileDiffusionDualGuidedPipeline",
841
+ "VersatileDiffusionImageVariationPipeline",
842
+ "VersatileDiffusionPipeline",
843
+ "VersatileDiffusionTextToImagePipeline",
844
+ "VideoToVideoSDPipeline",
845
+ "VisualClozeGenerationPipeline",
846
+ "VisualClozePipeline",
847
+ "VQDiffusionPipeline",
848
+ "WanAnimatePipeline",
849
+ "WanImageToVideoPipeline",
850
+ "WanPipeline",
851
+ "WanVACEPipeline",
852
+ "WanVideoToVideoPipeline",
853
+ "WuerstchenCombinedPipeline",
854
+ "WuerstchenDecoderPipeline",
855
+ "WuerstchenPriorPipeline",
856
+ "ZImageControlNetInpaintPipeline",
857
+ "ZImageControlNetPipeline",
858
+ "ZImageImg2ImgPipeline",
859
+ "ZImageInpaintPipeline",
860
+ "ZImageOmniPipeline",
861
+ "ZImagePipeline",
862
+ ]
863
+ )
864
+
865
+
866
+ try:
867
+ if not (is_torch_available() and is_transformers_available() and is_opencv_available()):
868
+ raise OptionalDependencyNotAvailable()
869
+ except OptionalDependencyNotAvailable:
870
+ from .utils import dummy_torch_and_transformers_and_opencv_objects # noqa F403
871
+
872
+ _import_structure["utils.dummy_torch_and_transformers_and_opencv_objects"] = [
873
+ name for name in dir(dummy_torch_and_transformers_and_opencv_objects) if not name.startswith("_")
874
+ ]
875
+
876
+ else:
877
+ _import_structure["pipelines"].extend(["ConsisIDPipeline"])
878
+
879
+ try:
880
+ if not (is_torch_available() and is_transformers_available() and is_sentencepiece_available()):
881
+ raise OptionalDependencyNotAvailable()
882
+ except OptionalDependencyNotAvailable:
883
+ from .utils import dummy_torch_and_transformers_and_sentencepiece_objects # noqa F403
884
+
885
+ _import_structure["utils.dummy_torch_and_transformers_and_sentencepiece_objects"] = [
886
+ name for name in dir(dummy_torch_and_transformers_and_sentencepiece_objects) if not name.startswith("_")
887
+ ]
888
+
889
+ else:
890
+ _import_structure["pipelines"].extend(["KolorsImg2ImgPipeline", "KolorsPAGPipeline", "KolorsPipeline"])
891
+
892
+ try:
893
+ if not (is_torch_available() and is_transformers_available() and is_onnx_available()):
894
+ raise OptionalDependencyNotAvailable()
895
+ except OptionalDependencyNotAvailable:
896
+ from .utils import dummy_torch_and_transformers_and_onnx_objects # noqa F403
897
+
898
+ _import_structure["utils.dummy_torch_and_transformers_and_onnx_objects"] = [
899
+ name for name in dir(dummy_torch_and_transformers_and_onnx_objects) if not name.startswith("_")
900
+ ]
901
+
902
+ else:
903
+ _import_structure["pipelines"].extend(
904
+ [
905
+ "OnnxStableDiffusionImg2ImgPipeline",
906
+ "OnnxStableDiffusionInpaintPipeline",
907
+ "OnnxStableDiffusionInpaintPipelineLegacy",
908
+ "OnnxStableDiffusionPipeline",
909
+ "OnnxStableDiffusionUpscalePipeline",
910
+ "StableDiffusionOnnxPipeline",
911
+ ]
912
+ )
913
+
914
+ try:
915
+ if not (is_torch_available() and is_librosa_available()):
916
+ raise OptionalDependencyNotAvailable()
917
+ except OptionalDependencyNotAvailable:
918
+ from .utils import dummy_torch_and_librosa_objects # noqa F403
919
+
920
+ _import_structure["utils.dummy_torch_and_librosa_objects"] = [
921
+ name for name in dir(dummy_torch_and_librosa_objects) if not name.startswith("_")
922
+ ]
923
+
924
+ else:
925
+ _import_structure["pipelines"].extend(["AudioDiffusionPipeline", "Mel"])
926
+
927
+ try:
928
+ if not (is_transformers_available() and is_torch_available() and is_note_seq_available()):
929
+ raise OptionalDependencyNotAvailable()
930
+ except OptionalDependencyNotAvailable:
931
+ from .utils import dummy_transformers_and_torch_and_note_seq_objects # noqa F403
932
+
933
+ _import_structure["utils.dummy_transformers_and_torch_and_note_seq_objects"] = [
934
+ name for name in dir(dummy_transformers_and_torch_and_note_seq_objects) if not name.startswith("_")
935
+ ]
936
+
937
+
938
+ else:
939
+ _import_structure["pipelines"].extend(["SpectrogramDiffusionPipeline"])
940
+
941
+ try:
942
+ if not (is_note_seq_available()):
943
+ raise OptionalDependencyNotAvailable()
944
+ except OptionalDependencyNotAvailable:
945
+ from .utils import dummy_note_seq_objects # noqa F403
946
+
947
+ _import_structure["utils.dummy_note_seq_objects"] = [
948
+ name for name in dir(dummy_note_seq_objects) if not name.startswith("_")
949
+ ]
950
+
951
+
952
+ else:
953
+ _import_structure["pipelines"].extend(["MidiProcessor"])
954
+
955
+ if TYPE_CHECKING or DIFFUSERS_SLOW_IMPORT:
956
+ from .configuration_utils import ConfigMixin
957
+ from .quantizers import PipelineQuantizationConfig
958
+
959
+ try:
960
+ if not is_bitsandbytes_available():
961
+ raise OptionalDependencyNotAvailable()
962
+ except OptionalDependencyNotAvailable:
963
+ from .utils.dummy_bitsandbytes_objects import *
964
+ else:
965
+ from .quantizers.quantization_config import BitsAndBytesConfig
966
+
967
+ try:
968
+ if not is_gguf_available():
969
+ raise OptionalDependencyNotAvailable()
970
+ except OptionalDependencyNotAvailable:
971
+ from .utils.dummy_gguf_objects import *
972
+ else:
973
+ from .quantizers.quantization_config import GGUFQuantizationConfig
974
+
975
+ try:
976
+ if not is_torchao_available():
977
+ raise OptionalDependencyNotAvailable()
978
+ except OptionalDependencyNotAvailable:
979
+ from .utils.dummy_torchao_objects import *
980
+ else:
981
+ from .quantizers.quantization_config import TorchAoConfig
982
+
983
+ try:
984
+ if not is_optimum_quanto_available():
985
+ raise OptionalDependencyNotAvailable()
986
+ except OptionalDependencyNotAvailable:
987
+ from .utils.dummy_optimum_quanto_objects import *
988
+ else:
989
+ from .quantizers.quantization_config import QuantoConfig
990
+
991
+ try:
992
+ if not is_nvidia_modelopt_available():
993
+ raise OptionalDependencyNotAvailable()
994
+ except OptionalDependencyNotAvailable:
995
+ from .utils.dummy_nvidia_modelopt_objects import *
996
+ else:
997
+ from .quantizers.quantization_config import NVIDIAModelOptConfig
998
+
999
+ try:
1000
+ if not is_torch_available():
1001
+ raise OptionalDependencyNotAvailable()
1002
+ except OptionalDependencyNotAvailable:
1003
+ from .utils.dummy_nunchaku_lite_objects import *
1004
+ else:
1005
+ from .quantizers.quantization_config import NunchakuLiteQuantizationConfig
1006
+
1007
+ try:
1008
+ if not is_auto_round_available():
1009
+ raise OptionalDependencyNotAvailable()
1010
+ except OptionalDependencyNotAvailable:
1011
+ from .utils.dummy_auto_round_objects import *
1012
+ else:
1013
+ from .quantizers.quantization_config import AutoRoundConfig
1014
+
1015
+ try:
1016
+ if not is_sdnq_available():
1017
+ raise OptionalDependencyNotAvailable()
1018
+ except OptionalDependencyNotAvailable:
1019
+ from .utils.dummy_sdnq_objects import *
1020
+ else:
1021
+ from .quantizers.quantization_config import SDNQConfig
1022
+
1023
+ try:
1024
+ if not is_onnx_available():
1025
+ raise OptionalDependencyNotAvailable()
1026
+ except OptionalDependencyNotAvailable:
1027
+ from .utils.dummy_onnx_objects import * # noqa F403
1028
+ else:
1029
+ from .pipelines import OnnxRuntimeModel
1030
+
1031
+ try:
1032
+ if not is_torch_available():
1033
+ raise OptionalDependencyNotAvailable()
1034
+ except OptionalDependencyNotAvailable:
1035
+ from .utils.dummy_pt_objects import * # noqa F403
1036
+ else:
1037
+ from .guiders import (
1038
+ AdaptiveProjectedGuidance,
1039
+ AdaptiveProjectedMixGuidance,
1040
+ AutoGuidance,
1041
+ BaseGuidance,
1042
+ ClassifierFreeGuidance,
1043
+ ClassifierFreeZeroStarGuidance,
1044
+ FrequencyDecoupledGuidance,
1045
+ PerturbedAttentionGuidance,
1046
+ SkipLayerGuidance,
1047
+ SmoothedEnergyGuidance,
1048
+ TangentialClassifierFreeGuidance,
1049
+ )
1050
+ from .hooks import (
1051
+ FasterCacheConfig,
1052
+ FirstBlockCacheConfig,
1053
+ HookRegistry,
1054
+ LayerSkipConfig,
1055
+ MagCacheConfig,
1056
+ PyramidAttentionBroadcastConfig,
1057
+ SmoothedEnergyGuidanceConfig,
1058
+ TaylorSeerCacheConfig,
1059
+ TextKVCacheConfig,
1060
+ apply_faster_cache,
1061
+ apply_first_block_cache,
1062
+ apply_layer_skip,
1063
+ apply_mag_cache,
1064
+ apply_pyramid_attention_broadcast,
1065
+ apply_taylorseer_cache,
1066
+ apply_text_kv_cache,
1067
+ )
1068
+ from .image_processor import (
1069
+ InpaintProcessor,
1070
+ IPAdapterMaskProcessor,
1071
+ PixArtImageProcessor,
1072
+ VaeImageProcessor,
1073
+ VaeImageProcessorLDM3D,
1074
+ )
1075
+ from .models import (
1076
+ AceStepTransformer1DModel,
1077
+ AllegroTransformer3DModel,
1078
+ AnimaTextConditioner,
1079
+ AnyFlowFARTransformer3DModel,
1080
+ AnyFlowTransformer3DModel,
1081
+ AsymmetricAutoencoderKL,
1082
+ AttentionBackendName,
1083
+ AuraFlowTransformer2DModel,
1084
+ AutoencoderDC,
1085
+ AutoencoderKL,
1086
+ AutoencoderKLAllegro,
1087
+ AutoencoderKLCogVideoX,
1088
+ AutoencoderKLCosmos,
1089
+ AutoencoderKLFlux2,
1090
+ AutoencoderKLHunyuanImage,
1091
+ AutoencoderKLHunyuanImageRefiner,
1092
+ AutoencoderKLHunyuanVideo,
1093
+ AutoencoderKLHunyuanVideo15,
1094
+ AutoencoderKLKVAE,
1095
+ AutoencoderKLKVAEVideo,
1096
+ AutoencoderKLLTX2Audio,
1097
+ AutoencoderKLLTX2Video,
1098
+ AutoencoderKLLTXVideo,
1099
+ AutoencoderKLMagvit,
1100
+ AutoencoderKLMiniMaxH3,
1101
+ AutoencoderKLMiniMaxH3Audio,
1102
+ AutoencoderKLMochi,
1103
+ AutoencoderKLQwenImage,
1104
+ AutoencoderKLTemporalDecoder,
1105
+ AutoencoderKLWan,
1106
+ AutoencoderOobleck,
1107
+ AutoencoderRAE,
1108
+ AutoencoderTiny,
1109
+ AutoencoderVidTok,
1110
+ AutoModel,
1111
+ BriaFiboTransformer2DModel,
1112
+ BriaTransformer2DModel,
1113
+ CacheMixin,
1114
+ ChromaTransformer2DModel,
1115
+ ChronoEditTransformer3DModel,
1116
+ CogVideoXTransformer3DModel,
1117
+ CogView3PlusTransformer2DModel,
1118
+ CogView4Transformer2DModel,
1119
+ ConsisIDTransformer3DModel,
1120
+ ConsistencyDecoderVAE,
1121
+ ContextParallelConfig,
1122
+ ControlNetModel,
1123
+ ControlNetUnionModel,
1124
+ ControlNetXSAdapter,
1125
+ Cosmos3AVAEAudioTokenizer,
1126
+ Cosmos3OmniTransformer,
1127
+ CosmosControlNetModel,
1128
+ CosmosTransformer3DModel,
1129
+ DiTTransformer2DModel,
1130
+ DreamLiteTransformer2DModel,
1131
+ DreamLiteUNetModel,
1132
+ EasyAnimateTransformer3DModel,
1133
+ ErnieImageTransformer2DModel,
1134
+ Flux2Transformer2DModel,
1135
+ FluxControlNetModel,
1136
+ FluxMultiControlNetModel,
1137
+ FluxTransformer2DModel,
1138
+ GlmImageTransformer2DModel,
1139
+ HeliosTransformer3DModel,
1140
+ HiDreamImageTransformer2DModel,
1141
+ HunyuanDiT2DControlNetModel,
1142
+ HunyuanDiT2DModel,
1143
+ HunyuanDiT2DMultiControlNetModel,
1144
+ HunyuanImageTransformer2DModel,
1145
+ HunyuanVideo15Transformer3DModel,
1146
+ HunyuanVideoFramepackTransformer3DModel,
1147
+ HunyuanVideoTransformer3DModel,
1148
+ I2VGenXLUNet,
1149
+ Ideogram4Transformer2DModel,
1150
+ JoyImageEditPlusTransformer3DModel,
1151
+ JoyImageEditTransformer3DModel,
1152
+ Kandinsky3UNet,
1153
+ Kandinsky5Transformer3DModel,
1154
+ Krea2Transformer2DModel,
1155
+ LatteTransformer3DModel,
1156
+ LongCatAudioDiTTransformer,
1157
+ LongCatAudioDiTVae,
1158
+ LongCatImageTransformer2DModel,
1159
+ LTX2VideoTransformer3DModel,
1160
+ LTXVideoTransformer3DModel,
1161
+ Lumina2Transformer2DModel,
1162
+ LuminaNextDiT2DModel,
1163
+ MiniMaxH3Transformer3DModel,
1164
+ MochiTransformer3DModel,
1165
+ ModelMixin,
1166
+ MotifVideoTransformer3DModel,
1167
+ MotionAdapter,
1168
+ MultiAdapter,
1169
+ MultiControlNetModel,
1170
+ NucleusMoEImageTransformer2DModel,
1171
+ OmniGenTransformer2DModel,
1172
+ OvisImageTransformer2DModel,
1173
+ ParallelConfig,
1174
+ PixArtTransformer2DModel,
1175
+ PriorTransformer,
1176
+ PRXTransformer2DModel,
1177
+ QwenImageControlNetModel,
1178
+ QwenImageMultiControlNetModel,
1179
+ QwenImageTransformer2DModel,
1180
+ SanaControlNetModel,
1181
+ SanaTransformer2DModel,
1182
+ SanaVideoTransformer3DModel,
1183
+ SD3ControlNetModel,
1184
+ SD3MultiControlNetModel,
1185
+ SD3Transformer2DModel,
1186
+ SkyReelsV2Transformer3DModel,
1187
+ SparseControlNetModel,
1188
+ StableAudioDiTModel,
1189
+ T2IAdapter,
1190
+ T5FilmDecoder,
1191
+ Transformer2DModel,
1192
+ TransformerTemporalModel,
1193
+ UNet1DModel,
1194
+ UNet2DConditionModel,
1195
+ UNet2DModel,
1196
+ UNet3DConditionModel,
1197
+ UNetControlNetXSModel,
1198
+ UNetMotionModel,
1199
+ UNetSpatioTemporalConditionModel,
1200
+ UVit2DModel,
1201
+ VQModel,
1202
+ WanAnimateTransformer3DModel,
1203
+ WanTransformer3DModel,
1204
+ WanVACETransformer3DModel,
1205
+ ZImageControlNetModel,
1206
+ ZImageTransformer2DModel,
1207
+ attention_backend,
1208
+ )
1209
+ from .modular_pipelines import (
1210
+ AutoPipelineBlocks,
1211
+ ComponentsManager,
1212
+ ComponentSpec,
1213
+ ConditionalPipelineBlocks,
1214
+ ConfigSpec,
1215
+ InputParam,
1216
+ LoopSequentialPipelineBlocks,
1217
+ ModularPipeline,
1218
+ ModularPipelineBlocks,
1219
+ OutputParam,
1220
+ SequentialPipelineBlocks,
1221
+ )
1222
+ from .optimization import (
1223
+ get_constant_schedule,
1224
+ get_constant_schedule_with_warmup,
1225
+ get_cosine_schedule_with_warmup,
1226
+ get_cosine_with_hard_restarts_schedule_with_warmup,
1227
+ get_linear_schedule_with_warmup,
1228
+ get_polynomial_decay_schedule_with_warmup,
1229
+ get_scheduler,
1230
+ )
1231
+ from .pipelines import (
1232
+ AudioPipelineOutput,
1233
+ AutoPipelineForImage2Image,
1234
+ AutoPipelineForInpainting,
1235
+ AutoPipelineForText2Audio,
1236
+ AutoPipelineForText2Image,
1237
+ BlipDiffusionControlNetPipeline,
1238
+ BlipDiffusionPipeline,
1239
+ CLIPImageProjection,
1240
+ ConsistencyModelPipeline,
1241
+ DanceDiffusionPipeline,
1242
+ DDIMPipeline,
1243
+ DDPMPipeline,
1244
+ DiffusionPipeline,
1245
+ DiTPipeline,
1246
+ ImagePipelineOutput,
1247
+ KarrasVePipeline,
1248
+ LDMPipeline,
1249
+ LDMSuperResolutionPipeline,
1250
+ PNDMPipeline,
1251
+ RePaintPipeline,
1252
+ ScoreSdeVePipeline,
1253
+ StableDiffusionMixin,
1254
+ )
1255
+ from .quantizers import DiffusersQuantizer
1256
+ from .schedulers import (
1257
+ AmusedScheduler,
1258
+ BlockRefinementScheduler,
1259
+ BlockRefinementSchedulerOutput,
1260
+ CMStochasticIterativeScheduler,
1261
+ CogVideoXDDIMScheduler,
1262
+ CogVideoXDPMScheduler,
1263
+ DDIMInverseScheduler,
1264
+ DDIMParallelScheduler,
1265
+ DDIMScheduler,
1266
+ DDPMParallelScheduler,
1267
+ DDPMScheduler,
1268
+ DDPMWuerstchenScheduler,
1269
+ DEISMultistepScheduler,
1270
+ DiscreteDDIMScheduler,
1271
+ DiscreteDDIMSchedulerOutput,
1272
+ DPMSolverMultistepInverseScheduler,
1273
+ DPMSolverMultistepScheduler,
1274
+ DPMSolverSinglestepScheduler,
1275
+ EDMDPMSolverMultistepScheduler,
1276
+ EDMEulerScheduler,
1277
+ EntropyBoundScheduler,
1278
+ EntropyBoundSchedulerOutput,
1279
+ EulerAncestralDiscreteScheduler,
1280
+ EulerDiscreteScheduler,
1281
+ FlowMapEulerDiscreteScheduler,
1282
+ FlowMatchEulerDiscreteScheduler,
1283
+ FlowMatchHeunDiscreteScheduler,
1284
+ FlowMatchLCMScheduler,
1285
+ HeliosDMDScheduler,
1286
+ HeliosScheduler,
1287
+ HeunDiscreteScheduler,
1288
+ IPNDMScheduler,
1289
+ KarrasVeScheduler,
1290
+ KDPM2AncestralDiscreteScheduler,
1291
+ KDPM2DiscreteScheduler,
1292
+ LCMScheduler,
1293
+ LTXEulerAncestralRFScheduler,
1294
+ MiniMaxH3Scheduler,
1295
+ PNDMScheduler,
1296
+ RePaintScheduler,
1297
+ SASolverScheduler,
1298
+ SchedulerMixin,
1299
+ SCMScheduler,
1300
+ ScoreSdeVeScheduler,
1301
+ TCDScheduler,
1302
+ UnCLIPScheduler,
1303
+ UniPCMultistepScheduler,
1304
+ VQDiffusionScheduler,
1305
+ )
1306
+ from .training_utils import EMAModel
1307
+ from .video_processor import VideoProcessor
1308
+
1309
+ try:
1310
+ if not (is_torch_available() and is_scipy_available()):
1311
+ raise OptionalDependencyNotAvailable()
1312
+ except OptionalDependencyNotAvailable:
1313
+ from .utils.dummy_torch_and_scipy_objects import * # noqa F403
1314
+ else:
1315
+ from .schedulers import LMSDiscreteScheduler
1316
+
1317
+ try:
1318
+ if not (is_torch_available() and is_torchsde_available()):
1319
+ raise OptionalDependencyNotAvailable()
1320
+ except OptionalDependencyNotAvailable:
1321
+ from .utils.dummy_torch_and_torchsde_objects import * # noqa F403
1322
+ else:
1323
+ from .schedulers import CosineDPMSolverMultistepScheduler, DPMSolverSDEScheduler
1324
+
1325
+ try:
1326
+ if not (is_torch_available() and is_transformers_available()):
1327
+ raise OptionalDependencyNotAvailable()
1328
+ except OptionalDependencyNotAvailable:
1329
+ from .utils.dummy_torch_and_transformers_objects import * # noqa F403
1330
+ else:
1331
+ from .modular_pipelines import (
1332
+ AnimaAutoBlocks,
1333
+ AnimaModularPipeline,
1334
+ Cosmos3DistilledBlocks,
1335
+ Cosmos3DistilledModularPipeline,
1336
+ Cosmos3OmniBlocks,
1337
+ Cosmos3OmniModularPipeline,
1338
+ ErnieImageAutoBlocks,
1339
+ ErnieImageModularPipeline,
1340
+ Flux2AutoBlocks,
1341
+ Flux2KleinAutoBlocks,
1342
+ Flux2KleinBaseAutoBlocks,
1343
+ Flux2KleinBaseModularPipeline,
1344
+ Flux2KleinModularPipeline,
1345
+ Flux2ModularPipeline,
1346
+ FluxAutoBlocks,
1347
+ FluxKontextAutoBlocks,
1348
+ FluxKontextModularPipeline,
1349
+ FluxModularPipeline,
1350
+ HeliosAutoBlocks,
1351
+ HeliosModularPipeline,
1352
+ HeliosPyramidAutoBlocks,
1353
+ HeliosPyramidDistilledAutoBlocks,
1354
+ HeliosPyramidDistilledModularPipeline,
1355
+ HeliosPyramidModularPipeline,
1356
+ HunyuanVideo15AutoBlocks,
1357
+ HunyuanVideo15ModularPipeline,
1358
+ Ideogram4AutoBlocks,
1359
+ Ideogram4ModularPipeline,
1360
+ Krea2AutoBlocks,
1361
+ Krea2ModularPipeline,
1362
+ Krea2TurboAutoBlocks,
1363
+ Krea2TurboModularPipeline,
1364
+ LTXAutoBlocks,
1365
+ LTXModularPipeline,
1366
+ MiniMaxH3Blocks,
1367
+ MiniMaxH3ModularPipeline,
1368
+ MiniMaxH3Ref2VABlocks,
1369
+ MiniMaxH3Ref2VAModularPipeline,
1370
+ QwenImageAutoBlocks,
1371
+ QwenImageEditAutoBlocks,
1372
+ QwenImageEditModularPipeline,
1373
+ QwenImageEditPlusAutoBlocks,
1374
+ QwenImageEditPlusModularPipeline,
1375
+ QwenImageLayeredAutoBlocks,
1376
+ QwenImageLayeredModularPipeline,
1377
+ QwenImageModularPipeline,
1378
+ StableDiffusion3AutoBlocks,
1379
+ StableDiffusion3ModularPipeline,
1380
+ StableDiffusionXLAutoBlocks,
1381
+ StableDiffusionXLModularPipeline,
1382
+ Wan22Blocks,
1383
+ Wan22Image2VideoBlocks,
1384
+ Wan22Image2VideoModularPipeline,
1385
+ Wan22ModularPipeline,
1386
+ WanBlocks,
1387
+ WanImage2VideoAutoBlocks,
1388
+ WanImage2VideoModularPipeline,
1389
+ WanModularPipeline,
1390
+ ZImageAutoBlocks,
1391
+ ZImageModularPipeline,
1392
+ )
1393
+ from .pipelines import (
1394
+ AceStepAudioTokenDetokenizer,
1395
+ AceStepAudioTokenizer,
1396
+ AceStepConditionEncoder,
1397
+ AceStepPipeline,
1398
+ AllegroPipeline,
1399
+ AltDiffusionImg2ImgPipeline,
1400
+ AltDiffusionPipeline,
1401
+ AmusedImg2ImgPipeline,
1402
+ AmusedInpaintPipeline,
1403
+ AmusedPipeline,
1404
+ AnimateDiffControlNetPipeline,
1405
+ AnimateDiffPAGPipeline,
1406
+ AnimateDiffPipeline,
1407
+ AnimateDiffSDXLPipeline,
1408
+ AnimateDiffSparseControlNetPipeline,
1409
+ AnimateDiffVideoToVideoControlNetPipeline,
1410
+ AnimateDiffVideoToVideoPipeline,
1411
+ AnyFlowFARPipeline,
1412
+ AnyFlowPipeline,
1413
+ AudioLDM2Pipeline,
1414
+ AudioLDM2ProjectionModel,
1415
+ AudioLDM2UNet2DConditionModel,
1416
+ AudioLDMPipeline,
1417
+ AuraFlowPipeline,
1418
+ BriaFiboEditPipeline,
1419
+ BriaFiboPipeline,
1420
+ BriaPipeline,
1421
+ ChromaImg2ImgPipeline,
1422
+ ChromaInpaintPipeline,
1423
+ ChromaPipeline,
1424
+ ChronoEditPipeline,
1425
+ CLIPImageProjection,
1426
+ CogVideoXFunControlPipeline,
1427
+ CogVideoXImageToVideoPipeline,
1428
+ CogVideoXPipeline,
1429
+ CogVideoXVideoToVideoPipeline,
1430
+ CogView3PlusPipeline,
1431
+ CogView4ControlPipeline,
1432
+ CogView4Pipeline,
1433
+ ConsisIDPipeline,
1434
+ Cosmos2_5_PredictBasePipeline,
1435
+ Cosmos2_5_TransferPipeline,
1436
+ Cosmos2TextToImagePipeline,
1437
+ Cosmos2VideoToWorldPipeline,
1438
+ Cosmos3OmniPipeline,
1439
+ CosmosActionCondition,
1440
+ CosmosTextToWorldPipeline,
1441
+ CosmosVideoToWorldPipeline,
1442
+ CycleDiffusionPipeline,
1443
+ DiffusionGemmaPipeline,
1444
+ DiffusionGemmaPipelineOutput,
1445
+ DreamLiteMobilePipeline,
1446
+ DreamLitePipeline,
1447
+ DreamLitePipelineOutput,
1448
+ EasyAnimateControlPipeline,
1449
+ EasyAnimateInpaintPipeline,
1450
+ EasyAnimatePipeline,
1451
+ ErnieImagePipeline,
1452
+ Flux2KleinInpaintPipeline,
1453
+ Flux2KleinKVPipeline,
1454
+ Flux2KleinPipeline,
1455
+ Flux2Pipeline,
1456
+ FluxControlImg2ImgPipeline,
1457
+ FluxControlInpaintPipeline,
1458
+ FluxControlNetImg2ImgPipeline,
1459
+ FluxControlNetInpaintPipeline,
1460
+ FluxControlNetPipeline,
1461
+ FluxControlPipeline,
1462
+ FluxFillPipeline,
1463
+ FluxImg2ImgPipeline,
1464
+ FluxInpaintPipeline,
1465
+ FluxKontextInpaintPipeline,
1466
+ FluxKontextPipeline,
1467
+ FluxPipeline,
1468
+ FluxPriorReduxPipeline,
1469
+ GlmImagePipeline,
1470
+ HeliosPipeline,
1471
+ HeliosPyramidPipeline,
1472
+ HiDreamImagePipeline,
1473
+ HunyuanDiTControlNetPipeline,
1474
+ HunyuanDiTPAGPipeline,
1475
+ HunyuanDiTPipeline,
1476
+ HunyuanImagePipeline,
1477
+ HunyuanImageRefinerPipeline,
1478
+ HunyuanSkyreelsImageToVideoPipeline,
1479
+ HunyuanVideo15ImageToVideoPipeline,
1480
+ HunyuanVideo15Pipeline,
1481
+ HunyuanVideoFramepackPipeline,
1482
+ HunyuanVideoImageToVideoPipeline,
1483
+ HunyuanVideoPipeline,
1484
+ I2VGenXLPipeline,
1485
+ Ideogram4Pipeline,
1486
+ Ideogram4PromptEnhancerHead,
1487
+ IFImg2ImgPipeline,
1488
+ IFImg2ImgSuperResolutionPipeline,
1489
+ IFInpaintingPipeline,
1490
+ IFInpaintingSuperResolutionPipeline,
1491
+ IFPipeline,
1492
+ IFSuperResolutionPipeline,
1493
+ ImageTextPipelineOutput,
1494
+ JoyImageEditPipeline,
1495
+ JoyImageEditPipelineOutput,
1496
+ JoyImageEditPlusPipeline,
1497
+ JoyImageEditPlusPipelineOutput,
1498
+ Kandinsky3Img2ImgPipeline,
1499
+ Kandinsky3Pipeline,
1500
+ Kandinsky5I2IPipeline,
1501
+ Kandinsky5I2VPipeline,
1502
+ Kandinsky5T2IPipeline,
1503
+ Kandinsky5T2VPipeline,
1504
+ KandinskyCombinedPipeline,
1505
+ KandinskyImg2ImgCombinedPipeline,
1506
+ KandinskyImg2ImgPipeline,
1507
+ KandinskyInpaintCombinedPipeline,
1508
+ KandinskyInpaintPipeline,
1509
+ KandinskyPipeline,
1510
+ KandinskyPriorPipeline,
1511
+ KandinskyV22CombinedPipeline,
1512
+ KandinskyV22ControlnetImg2ImgPipeline,
1513
+ KandinskyV22ControlnetPipeline,
1514
+ KandinskyV22Img2ImgCombinedPipeline,
1515
+ KandinskyV22Img2ImgPipeline,
1516
+ KandinskyV22InpaintCombinedPipeline,
1517
+ KandinskyV22InpaintPipeline,
1518
+ KandinskyV22Pipeline,
1519
+ KandinskyV22PriorEmb2EmbPipeline,
1520
+ KandinskyV22PriorPipeline,
1521
+ Krea2Pipeline,
1522
+ LatentConsistencyModelImg2ImgPipeline,
1523
+ LatentConsistencyModelPipeline,
1524
+ LattePipeline,
1525
+ LDMTextToImagePipeline,
1526
+ LEditsPPPipelineStableDiffusion,
1527
+ LEditsPPPipelineStableDiffusionXL,
1528
+ LLaDA2Pipeline,
1529
+ LLaDA2PipelineOutput,
1530
+ LongCatAudioDiTPipeline,
1531
+ LongCatImageEditPipeline,
1532
+ LongCatImagePipeline,
1533
+ LTX2ConditionPipeline,
1534
+ LTX2HDRPipeline,
1535
+ LTX2ImageToVideoPipeline,
1536
+ LTX2InContextPipeline,
1537
+ LTX2LatentUpsamplePipeline,
1538
+ LTX2Pipeline,
1539
+ LTXConditionPipeline,
1540
+ LTXI2VLongMultiPromptPipeline,
1541
+ LTXImageToVideoPipeline,
1542
+ LTXLatentUpsamplePipeline,
1543
+ LTXPipeline,
1544
+ LucyEditPipeline,
1545
+ Lumina2Pipeline,
1546
+ Lumina2Text2ImgPipeline,
1547
+ LuminaPipeline,
1548
+ LuminaText2ImgPipeline,
1549
+ MarigoldDepthPipeline,
1550
+ MarigoldIntrinsicsPipeline,
1551
+ MarigoldNormalsPipeline,
1552
+ MochiPipeline,
1553
+ MotifVideoImage2VideoPipeline,
1554
+ MotifVideoPipeline,
1555
+ MotifVideoPipelineOutput,
1556
+ MusicLDMPipeline,
1557
+ NucleusMoEImagePipeline,
1558
+ OmniGenPipeline,
1559
+ OvisImagePipeline,
1560
+ PaintByExamplePipeline,
1561
+ PIAPipeline,
1562
+ PixArtAlphaPipeline,
1563
+ PixArtSigmaPAGPipeline,
1564
+ PixArtSigmaPipeline,
1565
+ PRXPipeline,
1566
+ PRXPixelPipeline,
1567
+ QwenImageControlNetInpaintPipeline,
1568
+ QwenImageControlNetPipeline,
1569
+ QwenImageEditInpaintPipeline,
1570
+ QwenImageEditPipeline,
1571
+ QwenImageEditPlusPipeline,
1572
+ QwenImageImg2ImgPipeline,
1573
+ QwenImageInpaintPipeline,
1574
+ QwenImageLayeredPipeline,
1575
+ QwenImagePipeline,
1576
+ ReduxImageEncoder,
1577
+ SanaControlNetPipeline,
1578
+ SanaImageToVideoPipeline,
1579
+ SanaPAGPipeline,
1580
+ SanaPipeline,
1581
+ SanaSprintImg2ImgPipeline,
1582
+ SanaSprintPipeline,
1583
+ SanaVideoPipeline,
1584
+ SemanticStableDiffusionPipeline,
1585
+ ShapEImg2ImgPipeline,
1586
+ ShapEPipeline,
1587
+ SkyReelsV2DiffusionForcingImageToVideoPipeline,
1588
+ SkyReelsV2DiffusionForcingPipeline,
1589
+ SkyReelsV2DiffusionForcingVideoToVideoPipeline,
1590
+ SkyReelsV2ImageToVideoPipeline,
1591
+ SkyReelsV2Pipeline,
1592
+ StableAudioPipeline,
1593
+ StableAudioProjectionModel,
1594
+ StableCascadeCombinedPipeline,
1595
+ StableCascadeDecoderPipeline,
1596
+ StableCascadePriorPipeline,
1597
+ StableDiffusion3ControlNetInpaintingPipeline,
1598
+ StableDiffusion3ControlNetPipeline,
1599
+ StableDiffusion3Img2ImgPipeline,
1600
+ StableDiffusion3InpaintPipeline,
1601
+ StableDiffusion3PAGImg2ImgPipeline,
1602
+ StableDiffusion3PAGPipeline,
1603
+ StableDiffusion3Pipeline,
1604
+ StableDiffusionAdapterPipeline,
1605
+ StableDiffusionAttendAndExcitePipeline,
1606
+ StableDiffusionControlNetImg2ImgPipeline,
1607
+ StableDiffusionControlNetInpaintPipeline,
1608
+ StableDiffusionControlNetPAGInpaintPipeline,
1609
+ StableDiffusionControlNetPAGPipeline,
1610
+ StableDiffusionControlNetPipeline,
1611
+ StableDiffusionControlNetXSPipeline,
1612
+ StableDiffusionDepth2ImgPipeline,
1613
+ StableDiffusionDiffEditPipeline,
1614
+ StableDiffusionGLIGENPipeline,
1615
+ StableDiffusionGLIGENTextImagePipeline,
1616
+ StableDiffusionImageVariationPipeline,
1617
+ StableDiffusionImg2ImgPipeline,
1618
+ StableDiffusionInpaintPipeline,
1619
+ StableDiffusionInpaintPipelineLegacy,
1620
+ StableDiffusionInstructPix2PixPipeline,
1621
+ StableDiffusionLatentUpscalePipeline,
1622
+ StableDiffusionLDM3DPipeline,
1623
+ StableDiffusionModelEditingPipeline,
1624
+ StableDiffusionPAGImg2ImgPipeline,
1625
+ StableDiffusionPAGInpaintPipeline,
1626
+ StableDiffusionPAGPipeline,
1627
+ StableDiffusionPanoramaPipeline,
1628
+ StableDiffusionParadigmsPipeline,
1629
+ StableDiffusionPipeline,
1630
+ StableDiffusionPipelineSafe,
1631
+ StableDiffusionPix2PixZeroPipeline,
1632
+ StableDiffusionSAGPipeline,
1633
+ StableDiffusionUpscalePipeline,
1634
+ StableDiffusionXLAdapterPipeline,
1635
+ StableDiffusionXLControlNetImg2ImgPipeline,
1636
+ StableDiffusionXLControlNetInpaintPipeline,
1637
+ StableDiffusionXLControlNetPAGImg2ImgPipeline,
1638
+ StableDiffusionXLControlNetPAGPipeline,
1639
+ StableDiffusionXLControlNetPipeline,
1640
+ StableDiffusionXLControlNetUnionImg2ImgPipeline,
1641
+ StableDiffusionXLControlNetUnionInpaintPipeline,
1642
+ StableDiffusionXLControlNetUnionPipeline,
1643
+ StableDiffusionXLControlNetXSPipeline,
1644
+ StableDiffusionXLImg2ImgPipeline,
1645
+ StableDiffusionXLInpaintPipeline,
1646
+ StableDiffusionXLInstructPix2PixPipeline,
1647
+ StableDiffusionXLPAGImg2ImgPipeline,
1648
+ StableDiffusionXLPAGInpaintPipeline,
1649
+ StableDiffusionXLPAGPipeline,
1650
+ StableDiffusionXLPipeline,
1651
+ StableUnCLIPImg2ImgPipeline,
1652
+ StableUnCLIPPipeline,
1653
+ StableVideoDiffusionPipeline,
1654
+ TextToVideoSDPipeline,
1655
+ TextToVideoZeroPipeline,
1656
+ TextToVideoZeroSDXLPipeline,
1657
+ UnCLIPImageVariationPipeline,
1658
+ UnCLIPPipeline,
1659
+ UniDiffuserModel,
1660
+ UniDiffuserPipeline,
1661
+ UniDiffuserTextDecoder,
1662
+ VersatileDiffusionDualGuidedPipeline,
1663
+ VersatileDiffusionImageVariationPipeline,
1664
+ VersatileDiffusionPipeline,
1665
+ VersatileDiffusionTextToImagePipeline,
1666
+ VideoToVideoSDPipeline,
1667
+ VisualClozeGenerationPipeline,
1668
+ VisualClozePipeline,
1669
+ VQDiffusionPipeline,
1670
+ WanAnimatePipeline,
1671
+ WanImageToVideoPipeline,
1672
+ WanPipeline,
1673
+ WanVACEPipeline,
1674
+ WanVideoToVideoPipeline,
1675
+ WuerstchenCombinedPipeline,
1676
+ WuerstchenDecoderPipeline,
1677
+ WuerstchenPriorPipeline,
1678
+ ZImageControlNetInpaintPipeline,
1679
+ ZImageControlNetPipeline,
1680
+ ZImageImg2ImgPipeline,
1681
+ ZImageInpaintPipeline,
1682
+ ZImageOmniPipeline,
1683
+ ZImagePipeline,
1684
+ )
1685
+
1686
+ try:
1687
+ if not (is_torch_available() and is_transformers_available() and is_sentencepiece_available()):
1688
+ raise OptionalDependencyNotAvailable()
1689
+ except OptionalDependencyNotAvailable:
1690
+ from .utils.dummy_torch_and_transformers_and_sentencepiece_objects import * # noqa F403
1691
+ else:
1692
+ from .pipelines import KolorsImg2ImgPipeline, KolorsPAGPipeline, KolorsPipeline
1693
+
1694
+ try:
1695
+ if not (is_torch_available() and is_transformers_available() and is_opencv_available()):
1696
+ raise OptionalDependencyNotAvailable()
1697
+ except OptionalDependencyNotAvailable:
1698
+ from .utils.dummy_torch_and_transformers_and_opencv_objects import * # noqa F403
1699
+ else:
1700
+ from .pipelines import ConsisIDPipeline
1701
+
1702
+ try:
1703
+ if not (is_torch_available() and is_transformers_available() and is_onnx_available()):
1704
+ raise OptionalDependencyNotAvailable()
1705
+ except OptionalDependencyNotAvailable:
1706
+ from .utils.dummy_torch_and_transformers_and_onnx_objects import * # noqa F403
1707
+ else:
1708
+ from .pipelines import (
1709
+ OnnxStableDiffusionImg2ImgPipeline,
1710
+ OnnxStableDiffusionInpaintPipeline,
1711
+ OnnxStableDiffusionInpaintPipelineLegacy,
1712
+ OnnxStableDiffusionPipeline,
1713
+ OnnxStableDiffusionUpscalePipeline,
1714
+ StableDiffusionOnnxPipeline,
1715
+ )
1716
+
1717
+ try:
1718
+ if not (is_torch_available() and is_librosa_available()):
1719
+ raise OptionalDependencyNotAvailable()
1720
+ except OptionalDependencyNotAvailable:
1721
+ from .utils.dummy_torch_and_librosa_objects import * # noqa F403
1722
+ else:
1723
+ from .pipelines import AudioDiffusionPipeline, Mel
1724
+
1725
+ try:
1726
+ if not (is_transformers_available() and is_torch_available() and is_note_seq_available()):
1727
+ raise OptionalDependencyNotAvailable()
1728
+ except OptionalDependencyNotAvailable:
1729
+ from .utils.dummy_transformers_and_torch_and_note_seq_objects import * # noqa F403
1730
+ else:
1731
+ from .pipelines import SpectrogramDiffusionPipeline
1732
+
1733
+ try:
1734
+ if not (is_note_seq_available()):
1735
+ raise OptionalDependencyNotAvailable()
1736
+ except OptionalDependencyNotAvailable:
1737
+ from .utils.dummy_note_seq_objects import * # noqa F403
1738
+ else:
1739
+ from .pipelines import MidiProcessor
1740
+
1741
+ else:
1742
+ import sys
1743
+
1744
+ sys.modules[__name__] = _LazyModule(
1745
+ __name__,
1746
+ globals()["__file__"],
1747
+ _import_structure,
1748
+ module_spec=__spec__,
1749
+ extra_objects={"__version__": __version__},
1750
+ )
diffusers/callbacks.py ADDED
@@ -0,0 +1,244 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from typing import Any
2
+
3
+ from .configuration_utils import ConfigMixin, register_to_config
4
+ from .utils import CONFIG_NAME
5
+
6
+
7
+ class PipelineCallback(ConfigMixin):
8
+ """
9
+ Base class for all the official callbacks used in a pipeline. This class provides a structure for implementing
10
+ custom callbacks and ensures that all callbacks have a consistent interface.
11
+
12
+ Please implement the following:
13
+ `tensor_inputs`: This should return a list of tensor inputs specific to your callback. You will only be able to
14
+ include
15
+ variables listed in the `._callback_tensor_inputs` attribute of your pipeline class.
16
+ `callback_fn`: This method defines the core functionality of your callback.
17
+ """
18
+
19
+ config_name = CONFIG_NAME
20
+
21
+ @register_to_config
22
+ def __init__(self, cutoff_step_ratio=1.0, cutoff_step_index=None):
23
+ super().__init__()
24
+
25
+ if (cutoff_step_ratio is None and cutoff_step_index is None) or (
26
+ cutoff_step_ratio is not None and cutoff_step_index is not None
27
+ ):
28
+ raise ValueError("Either cutoff_step_ratio or cutoff_step_index should be provided, not both or none.")
29
+
30
+ if cutoff_step_ratio is not None and (
31
+ not isinstance(cutoff_step_ratio, float) or not (0.0 <= cutoff_step_ratio <= 1.0)
32
+ ):
33
+ raise ValueError("cutoff_step_ratio must be a float between 0.0 and 1.0.")
34
+
35
+ @property
36
+ def tensor_inputs(self) -> list[str]:
37
+ raise NotImplementedError(f"You need to set the attribute `tensor_inputs` for {self.__class__}")
38
+
39
+ def callback_fn(self, pipeline, step_index, timesteps, callback_kwargs) -> dict[str, Any]:
40
+ raise NotImplementedError(f"You need to implement the method `callback_fn` for {self.__class__}")
41
+
42
+ def __call__(self, pipeline, step_index, timestep, callback_kwargs) -> dict[str, Any]:
43
+ return self.callback_fn(pipeline, step_index, timestep, callback_kwargs)
44
+
45
+
46
+ class MultiPipelineCallbacks:
47
+ """
48
+ This class is designed to handle multiple pipeline callbacks. It accepts a list of PipelineCallback objects and
49
+ provides a unified interface for calling all of them.
50
+ """
51
+
52
+ def __init__(self, callbacks: list[PipelineCallback]):
53
+ self.callbacks = callbacks
54
+
55
+ @property
56
+ def tensor_inputs(self) -> list[str]:
57
+ return [input for callback in self.callbacks for input in callback.tensor_inputs]
58
+
59
+ def __call__(self, pipeline, step_index, timestep, callback_kwargs) -> dict[str, Any]:
60
+ """
61
+ Calls all the callbacks in order with the given arguments and returns the final callback_kwargs.
62
+ """
63
+ for callback in self.callbacks:
64
+ callback_kwargs = callback(pipeline, step_index, timestep, callback_kwargs)
65
+
66
+ return callback_kwargs
67
+
68
+
69
+ class SDCFGCutoffCallback(PipelineCallback):
70
+ """
71
+ Callback function for Stable Diffusion Pipelines. After certain number of steps (set by `cutoff_step_ratio` or
72
+ `cutoff_step_index`), this callback will disable the CFG.
73
+
74
+ Note: This callback mutates the pipeline by changing the `_guidance_scale` attribute to 0.0 after the cutoff step.
75
+ """
76
+
77
+ tensor_inputs = ["prompt_embeds"]
78
+
79
+ def callback_fn(self, pipeline, step_index, timestep, callback_kwargs) -> dict[str, Any]:
80
+ cutoff_step_ratio = self.config.cutoff_step_ratio
81
+ cutoff_step_index = self.config.cutoff_step_index
82
+
83
+ # Use cutoff_step_index if it's not None, otherwise use cutoff_step_ratio
84
+ cutoff_step = (
85
+ cutoff_step_index if cutoff_step_index is not None else int(pipeline.num_timesteps * cutoff_step_ratio)
86
+ )
87
+
88
+ if step_index == cutoff_step:
89
+ prompt_embeds = callback_kwargs[self.tensor_inputs[0]]
90
+ prompt_embeds = prompt_embeds[-1:] # "-1" denotes the embeddings for conditional text tokens.
91
+
92
+ pipeline._guidance_scale = 0.0
93
+
94
+ callback_kwargs[self.tensor_inputs[0]] = prompt_embeds
95
+ return callback_kwargs
96
+
97
+
98
+ class SDXLCFGCutoffCallback(PipelineCallback):
99
+ """
100
+ Callback function for the base Stable Diffusion XL Pipelines. After certain number of steps (set by
101
+ `cutoff_step_ratio` or `cutoff_step_index`), this callback will disable the CFG.
102
+
103
+ Note: This callback mutates the pipeline by changing the `_guidance_scale` attribute to 0.0 after the cutoff step.
104
+ """
105
+
106
+ tensor_inputs = [
107
+ "prompt_embeds",
108
+ "add_text_embeds",
109
+ "add_time_ids",
110
+ ]
111
+
112
+ def callback_fn(self, pipeline, step_index, timestep, callback_kwargs) -> dict[str, Any]:
113
+ cutoff_step_ratio = self.config.cutoff_step_ratio
114
+ cutoff_step_index = self.config.cutoff_step_index
115
+
116
+ # Use cutoff_step_index if it's not None, otherwise use cutoff_step_ratio
117
+ cutoff_step = (
118
+ cutoff_step_index if cutoff_step_index is not None else int(pipeline.num_timesteps * cutoff_step_ratio)
119
+ )
120
+
121
+ if step_index == cutoff_step:
122
+ prompt_embeds = callback_kwargs[self.tensor_inputs[0]]
123
+ prompt_embeds = prompt_embeds[-1:] # "-1" denotes the embeddings for conditional text tokens.
124
+
125
+ add_text_embeds = callback_kwargs[self.tensor_inputs[1]]
126
+ add_text_embeds = add_text_embeds[-1:] # "-1" denotes the embeddings for conditional pooled text tokens
127
+
128
+ add_time_ids = callback_kwargs[self.tensor_inputs[2]]
129
+ add_time_ids = add_time_ids[-1:] # "-1" denotes the embeddings for conditional added time vector
130
+
131
+ pipeline._guidance_scale = 0.0
132
+
133
+ callback_kwargs[self.tensor_inputs[0]] = prompt_embeds
134
+ callback_kwargs[self.tensor_inputs[1]] = add_text_embeds
135
+ callback_kwargs[self.tensor_inputs[2]] = add_time_ids
136
+
137
+ return callback_kwargs
138
+
139
+
140
+ class SDXLControlnetCFGCutoffCallback(PipelineCallback):
141
+ """
142
+ Callback function for the Controlnet Stable Diffusion XL Pipelines. After certain number of steps (set by
143
+ `cutoff_step_ratio` or `cutoff_step_index`), this callback will disable the CFG.
144
+
145
+ Note: This callback mutates the pipeline by changing the `_guidance_scale` attribute to 0.0 after the cutoff step.
146
+ """
147
+
148
+ tensor_inputs = [
149
+ "prompt_embeds",
150
+ "add_text_embeds",
151
+ "add_time_ids",
152
+ "image",
153
+ ]
154
+
155
+ def callback_fn(self, pipeline, step_index, timestep, callback_kwargs) -> dict[str, Any]:
156
+ cutoff_step_ratio = self.config.cutoff_step_ratio
157
+ cutoff_step_index = self.config.cutoff_step_index
158
+
159
+ # Use cutoff_step_index if it's not None, otherwise use cutoff_step_ratio
160
+ cutoff_step = (
161
+ cutoff_step_index if cutoff_step_index is not None else int(pipeline.num_timesteps * cutoff_step_ratio)
162
+ )
163
+
164
+ if step_index == cutoff_step:
165
+ prompt_embeds = callback_kwargs[self.tensor_inputs[0]]
166
+ prompt_embeds = prompt_embeds[-1:] # "-1" denotes the embeddings for conditional text tokens.
167
+
168
+ add_text_embeds = callback_kwargs[self.tensor_inputs[1]]
169
+ add_text_embeds = add_text_embeds[-1:] # "-1" denotes the embeddings for conditional pooled text tokens
170
+
171
+ add_time_ids = callback_kwargs[self.tensor_inputs[2]]
172
+ add_time_ids = add_time_ids[-1:] # "-1" denotes the embeddings for conditional added time vector
173
+
174
+ # For Controlnet
175
+ image = callback_kwargs[self.tensor_inputs[3]]
176
+ image = image[-1:]
177
+
178
+ pipeline._guidance_scale = 0.0
179
+
180
+ callback_kwargs[self.tensor_inputs[0]] = prompt_embeds
181
+ callback_kwargs[self.tensor_inputs[1]] = add_text_embeds
182
+ callback_kwargs[self.tensor_inputs[2]] = add_time_ids
183
+ callback_kwargs[self.tensor_inputs[3]] = image
184
+
185
+ return callback_kwargs
186
+
187
+
188
+ class IPAdapterScaleCutoffCallback(PipelineCallback):
189
+ """
190
+ Callback function for any pipeline that inherits `IPAdapterMixin`. After certain number of steps (set by
191
+ `cutoff_step_ratio` or `cutoff_step_index`), this callback will set the IP Adapter scale to `0.0`.
192
+
193
+ Note: This callback mutates the IP Adapter attention processors by setting the scale to 0.0 after the cutoff step.
194
+ """
195
+
196
+ tensor_inputs = []
197
+
198
+ def callback_fn(self, pipeline, step_index, timestep, callback_kwargs) -> dict[str, Any]:
199
+ cutoff_step_ratio = self.config.cutoff_step_ratio
200
+ cutoff_step_index = self.config.cutoff_step_index
201
+
202
+ # Use cutoff_step_index if it's not None, otherwise use cutoff_step_ratio
203
+ cutoff_step = (
204
+ cutoff_step_index if cutoff_step_index is not None else int(pipeline.num_timesteps * cutoff_step_ratio)
205
+ )
206
+
207
+ if step_index == cutoff_step:
208
+ pipeline.set_ip_adapter_scale(0.0)
209
+ return callback_kwargs
210
+
211
+
212
+ class SD3CFGCutoffCallback(PipelineCallback):
213
+ """
214
+ Callback function for Stable Diffusion 3 Pipelines. After certain number of steps (set by `cutoff_step_ratio` or
215
+ `cutoff_step_index`), this callback will disable the CFG.
216
+
217
+ Note: This callback mutates the pipeline by changing the `_guidance_scale` attribute to 0.0 after the cutoff step.
218
+ """
219
+
220
+ tensor_inputs = ["prompt_embeds", "pooled_prompt_embeds"]
221
+
222
+ def callback_fn(self, pipeline, step_index, timestep, callback_kwargs) -> dict[str, Any]:
223
+ cutoff_step_ratio = self.config.cutoff_step_ratio
224
+ cutoff_step_index = self.config.cutoff_step_index
225
+
226
+ # Use cutoff_step_index if it's not None, otherwise use cutoff_step_ratio
227
+ cutoff_step = (
228
+ cutoff_step_index if cutoff_step_index is not None else int(pipeline.num_timesteps * cutoff_step_ratio)
229
+ )
230
+
231
+ if step_index == cutoff_step:
232
+ prompt_embeds = callback_kwargs[self.tensor_inputs[0]]
233
+ prompt_embeds = prompt_embeds[-1:] # "-1" denotes the embeddings for conditional text tokens.
234
+
235
+ pooled_prompt_embeds = callback_kwargs[self.tensor_inputs[1]]
236
+ pooled_prompt_embeds = pooled_prompt_embeds[
237
+ -1:
238
+ ] # "-1" denotes the embeddings for conditional pooled text tokens.
239
+
240
+ pipeline._guidance_scale = 0.0
241
+
242
+ callback_kwargs[self.tensor_inputs[0]] = prompt_embeds
243
+ callback_kwargs[self.tensor_inputs[1]] = pooled_prompt_embeds
244
+ return callback_kwargs
diffusers/commands/__init__.py ADDED
@@ -0,0 +1,27 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 The HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ from abc import ABC, abstractmethod
16
+ from argparse import ArgumentParser
17
+
18
+
19
+ class BaseDiffusersCLICommand(ABC):
20
+ @staticmethod
21
+ @abstractmethod
22
+ def register_subcommand(parser: ArgumentParser):
23
+ raise NotImplementedError()
24
+
25
+ @abstractmethod
26
+ def run(self):
27
+ raise NotImplementedError()
diffusers/commands/custom_blocks.py ADDED
@@ -0,0 +1,140 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 The HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ """`diffusers-cli custom_blocks` — package a local `ModularPipelineBlocks` subclass for the Hub.
16
+
17
+ Parses `block.py` (or `--block_module_name`), instantiates the chosen block, and calls `save_pretrained` in the current
18
+ working directory.
19
+ """
20
+
21
+ import ast
22
+ import importlib.util
23
+ import os
24
+ from argparse import ArgumentParser, Namespace
25
+ from pathlib import Path
26
+
27
+ from ..utils import logging
28
+ from . import BaseDiffusersCLICommand
29
+
30
+
31
+ EXPECTED_PARENT_CLASSES = ["ModularPipelineBlocks"]
32
+
33
+
34
+ def conversion_command_factory(args: Namespace):
35
+ return CustomBlocksCommand(args.block_module_name, args.block_class_name)
36
+
37
+
38
+ class CustomBlocksCommand(BaseDiffusersCLICommand):
39
+ @staticmethod
40
+ def register_subcommand(parser: ArgumentParser):
41
+ from argparse import RawDescriptionHelpFormatter
42
+
43
+ epilog = (
44
+ "Examples\n"
45
+ " $ diffusers-cli custom_blocks\n"
46
+ " $ diffusers-cli custom_blocks --block_module_name my_block.py\n"
47
+ " $ diffusers-cli custom_blocks --block_module_name my_block.py --block_class_name MyDenoiseBlock\n"
48
+ "\n"
49
+ "Learn more\n"
50
+ " Use `diffusers-cli <command> --help` for more information about a command.\n"
51
+ " Read the documentation at https://huggingface.co/docs/diffusers\n"
52
+ )
53
+
54
+ conversion_parser = parser.add_parser(
55
+ "custom_blocks",
56
+ help="Package a local ModularPipelineBlocks subclass for the Hub.",
57
+ usage="\n diffusers-cli custom_blocks [options]",
58
+ epilog=epilog,
59
+ formatter_class=RawDescriptionHelpFormatter,
60
+ )
61
+ conversion_parser._optionals.title = "Options"
62
+ conversion_parser.add_argument(
63
+ "--block_module_name",
64
+ type=str,
65
+ default="block.py",
66
+ help="Module filename in which the custom block will be implemented.",
67
+ )
68
+ conversion_parser.add_argument(
69
+ "--block_class_name",
70
+ type=str,
71
+ default=None,
72
+ help="Name of the custom block. If provided None, we will try to infer it.",
73
+ )
74
+ conversion_parser.set_defaults(func=conversion_command_factory)
75
+
76
+ def __init__(self, block_module_name: str = "block.py", block_class_name: str = None):
77
+ self.logger = logging.get_logger("diffusers-cli/custom_blocks")
78
+ self.block_module_name = Path(block_module_name)
79
+ self.block_class_name = block_class_name
80
+
81
+ def run(self):
82
+ # determine the block to be saved.
83
+ out = self._get_class_names(self.block_module_name)
84
+ classes_found = list({cls for cls, _ in out})
85
+
86
+ if self.block_class_name is not None:
87
+ child_class, parent_class = self._choose_block(out, self.block_class_name)
88
+ if child_class is None and parent_class is None:
89
+ raise ValueError(
90
+ "`block_class_name` could not be retrieved. Available classes from "
91
+ f"{self.block_module_name}:\n{classes_found}"
92
+ )
93
+ else:
94
+ self.logger.info(
95
+ f"Found classes: {classes_found} will be using {classes_found[0]}. "
96
+ "If this needs to be changed, re-run the command specifying `block_class_name`."
97
+ )
98
+ child_class, parent_class = out[0][0], out[0][1]
99
+
100
+ # dynamically get the custom block and initialize it to call `save_pretrained` in the current directory.
101
+ # the user is responsible for running it, so I guess that is safe?
102
+ module_name = f"__dynamic__{self.block_module_name.stem}"
103
+ spec = importlib.util.spec_from_file_location(module_name, str(self.block_module_name))
104
+ module = importlib.util.module_from_spec(spec)
105
+ spec.loader.exec_module(module)
106
+ getattr(module, child_class)().save_pretrained(os.getcwd())
107
+
108
+ def _choose_block(self, candidates, chosen=None):
109
+ for cls, base in candidates:
110
+ if cls == chosen:
111
+ return cls, base
112
+ return None, None
113
+
114
+ def _get_class_names(self, file_path):
115
+ source = file_path.read_text(encoding="utf-8")
116
+ try:
117
+ tree = ast.parse(source, filename=file_path)
118
+ except SyntaxError as e:
119
+ raise ValueError(f"Could not parse {file_path!r}: {e}") from e
120
+
121
+ results: list[tuple[str, str]] = []
122
+ for node in tree.body:
123
+ if not isinstance(node, ast.ClassDef):
124
+ continue
125
+
126
+ base_names = [bname for b in node.bases if (bname := self._get_base_name(b)) is not None]
127
+
128
+ for allowed in EXPECTED_PARENT_CLASSES:
129
+ if allowed in base_names:
130
+ results.append((node.name, allowed))
131
+
132
+ return results
133
+
134
+ def _get_base_name(self, node: ast.expr):
135
+ if isinstance(node, ast.Name):
136
+ return node.id
137
+ elif isinstance(node, ast.Attribute):
138
+ val = self._get_base_name(node.value)
139
+ return f"{val}.{node.attr}" if val else node.attr
140
+ return None
diffusers/commands/diffusers_cli.py ADDED
@@ -0,0 +1,69 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python
2
+ # Copyright 2026 The HuggingFace Team. All rights reserved.
3
+ #
4
+ # Licensed under the Apache License, Version 2.0 (the "License");
5
+ # you may not use this file except in compliance with the License.
6
+ # You may obtain a copy of the License at
7
+ #
8
+ # http://www.apache.org/licenses/LICENSE-2.0
9
+ #
10
+ # Unless required by applicable law or agreed to in writing, software
11
+ # distributed under the License is distributed on an "AS IS" BASIS,
12
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ # See the License for the specific language governing permissions and
14
+ # limitations under the License.
15
+
16
+ from argparse import ArgumentParser
17
+
18
+ from huggingface_hub.cli._output import OutputFormat, out
19
+
20
+ from .custom_blocks import CustomBlocksCommand
21
+ from .env import EnvironmentCommand
22
+ from .fp16_safetensors import FP16SafetensorsCommand
23
+ from .run import RunCommand
24
+ from .schema import SchemaCommand
25
+ from .skills import SkillsCommand
26
+
27
+
28
+ def main():
29
+ parser = ArgumentParser(
30
+ prog="diffusers-cli",
31
+ usage="\n diffusers-cli [--format <fmt>] <command> [options]",
32
+ )
33
+ parser._optionals.title = "Options"
34
+ parser.add_argument(
35
+ "--format",
36
+ choices=[m.value for m in OutputFormat],
37
+ default=OutputFormat.auto.value,
38
+ help=(
39
+ "Output format. 'auto' (default) picks 'agent' when an AI coding agent is detected "
40
+ "(via CLAUDECODE/CURSOR_AI/AIDER_AI_CONTEXT/... env vars) and 'human' otherwise. "
41
+ "Must appear before the subcommand."
42
+ ),
43
+ )
44
+ commands_parser = parser.add_subparsers(title="Commands", metavar="<command>")
45
+
46
+ # Register commands
47
+ EnvironmentCommand.register_subcommand(commands_parser)
48
+ FP16SafetensorsCommand.register_subcommand(commands_parser)
49
+ CustomBlocksCommand.register_subcommand(commands_parser)
50
+ RunCommand.register_subcommand(commands_parser)
51
+ SchemaCommand.register_subcommand(commands_parser)
52
+ SkillsCommand.register_subcommand(commands_parser)
53
+
54
+ # Let's go
55
+ args = parser.parse_args()
56
+
57
+ out.set_mode(OutputFormat(args.format))
58
+
59
+ if not hasattr(args, "func"):
60
+ parser.print_help()
61
+ exit(1)
62
+
63
+ # Run
64
+ service = args.func(args)
65
+ service.run()
66
+
67
+
68
+ if __name__ == "__main__":
69
+ main()
diffusers/commands/env.py ADDED
@@ -0,0 +1,185 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 The HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ import importlib.metadata
16
+ import platform
17
+ import subprocess
18
+ from argparse import ArgumentParser
19
+
20
+ import huggingface_hub
21
+
22
+ from .. import __version__ as version
23
+ from ..utils import (
24
+ is_accelerate_available,
25
+ is_bitsandbytes_available,
26
+ is_gguf_available,
27
+ is_google_colab,
28
+ is_nvidia_modelopt_available,
29
+ is_optimum_quanto_available,
30
+ is_peft_available,
31
+ is_safetensors_available,
32
+ is_torch_available,
33
+ is_torchao_available,
34
+ is_transformers_available,
35
+ is_xformers_available,
36
+ )
37
+ from . import BaseDiffusersCLICommand
38
+
39
+
40
+ # (display name, availability_fn, pypi distribution name for importlib.metadata.version)
41
+ _QUANTIZATION_BACKENDS = (
42
+ ("bitsandbytes", is_bitsandbytes_available, "bitsandbytes"),
43
+ ("gguf", is_gguf_available, "gguf"),
44
+ ("optimum-quanto", is_optimum_quanto_available, "optimum-quanto"),
45
+ ("torchao", is_torchao_available, "torchao"),
46
+ ("nvidia-modelopt", is_nvidia_modelopt_available, "nvidia-modelopt"),
47
+ )
48
+
49
+
50
+ def info_command_factory(_):
51
+ return EnvironmentCommand()
52
+
53
+
54
+ class EnvironmentCommand(BaseDiffusersCLICommand):
55
+ @staticmethod
56
+ def register_subcommand(parser: ArgumentParser) -> None:
57
+ download_parser = parser.add_parser(
58
+ "env",
59
+ help="Print versions of diffusers and its dependencies (for bug reports).",
60
+ usage="\n diffusers-cli env",
61
+ )
62
+ download_parser._optionals.title = "Options"
63
+ download_parser.set_defaults(func=info_command_factory)
64
+
65
+ def run(self) -> dict:
66
+ hub_version = huggingface_hub.__version__
67
+
68
+ safetensors_version = "not installed"
69
+ if is_safetensors_available():
70
+ import safetensors
71
+
72
+ safetensors_version = safetensors.__version__
73
+
74
+ pt_version = "not installed"
75
+ pt_cuda_available = "NA"
76
+ if is_torch_available():
77
+ import torch
78
+
79
+ pt_version = torch.__version__
80
+ pt_cuda_available = torch.cuda.is_available()
81
+
82
+ transformers_version = "not installed"
83
+ if is_transformers_available():
84
+ import transformers
85
+
86
+ transformers_version = transformers.__version__
87
+
88
+ accelerate_version = "not installed"
89
+ if is_accelerate_available():
90
+ import accelerate
91
+
92
+ accelerate_version = accelerate.__version__
93
+
94
+ peft_version = "not installed"
95
+ if is_peft_available():
96
+ import peft
97
+
98
+ peft_version = peft.__version__
99
+
100
+ quantization_versions = {}
101
+ for backend_name, is_available_fn, dist_name in _QUANTIZATION_BACKENDS:
102
+ if not is_available_fn():
103
+ continue
104
+ try:
105
+ quantization_versions[backend_name] = importlib.metadata.version(dist_name)
106
+ except importlib.metadata.PackageNotFoundError:
107
+ quantization_versions[backend_name] = "N/A"
108
+
109
+ xformers_version = "not installed"
110
+ if is_xformers_available():
111
+ import xformers
112
+
113
+ xformers_version = xformers.__version__
114
+
115
+ platform_info = platform.platform()
116
+
117
+ is_google_colab_str = "Yes" if is_google_colab() else "No"
118
+
119
+ accelerator = "NA"
120
+ if platform.system() in {"Linux", "Windows"}:
121
+ try:
122
+ sp = subprocess.Popen(
123
+ ["nvidia-smi", "--query-gpu=gpu_name,memory.total", "--format=csv,noheader"],
124
+ stdout=subprocess.PIPE,
125
+ stderr=subprocess.PIPE,
126
+ )
127
+ out_str, _ = sp.communicate()
128
+ out_str = out_str.decode("utf-8")
129
+
130
+ if len(out_str) > 0:
131
+ accelerator = out_str.strip()
132
+ except FileNotFoundError:
133
+ pass
134
+ elif platform.system() == "Darwin": # Mac OS
135
+ try:
136
+ sp = subprocess.Popen(
137
+ ["system_profiler", "SPDisplaysDataType"],
138
+ stdout=subprocess.PIPE,
139
+ stderr=subprocess.PIPE,
140
+ )
141
+ out_str, _ = sp.communicate()
142
+ out_str = out_str.decode("utf-8")
143
+
144
+ start = out_str.find("Chipset Model:")
145
+ if start != -1:
146
+ start += len("Chipset Model:")
147
+ end = out_str.find("\n", start)
148
+ accelerator = out_str[start:end].strip()
149
+
150
+ start = out_str.find("VRAM (Total):")
151
+ if start != -1:
152
+ start += len("VRAM (Total):")
153
+ end = out_str.find("\n", start)
154
+ accelerator += " VRAM: " + out_str[start:end].strip()
155
+ except FileNotFoundError:
156
+ pass
157
+ else:
158
+ print("It seems you are running an unusual OS. Could you fill in the accelerator manually?")
159
+
160
+ info = {
161
+ "🤗 Diffusers version": version,
162
+ "Platform": platform_info,
163
+ "Running on Google Colab?": is_google_colab_str,
164
+ "Python version": platform.python_version(),
165
+ "PyTorch version (GPU?)": f"{pt_version} ({pt_cuda_available})",
166
+ "Huggingface_hub version": hub_version,
167
+ "Transformers version": transformers_version,
168
+ "Accelerate version": accelerate_version,
169
+ "PEFT version": peft_version,
170
+ **{f"{name} version": ver for name, ver in quantization_versions.items()},
171
+ "Safetensors version": safetensors_version,
172
+ "xFormers version": xformers_version,
173
+ "Accelerator": accelerator,
174
+ "Using GPU in script?": "<fill in>",
175
+ "Using distributed or parallel set-up in script?": "<fill in>",
176
+ }
177
+
178
+ print("\nCopy-and-paste the text below in your GitHub issue and FILL OUT the two last points.\n")
179
+ print(self.format_dict(info))
180
+
181
+ return info
182
+
183
+ @staticmethod
184
+ def format_dict(d: dict) -> str:
185
+ return "\n".join([f"- {prop}: {val}" for prop, val in d.items()]) + "\n"
diffusers/commands/fp16_safetensors.py ADDED
@@ -0,0 +1,144 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 The HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ """
16
+ Usage example:
17
+ diffusers-cli fp16_safetensors --ckpt_id=openai/shap-e --fp16 --use_safetensors
18
+ """
19
+
20
+ import glob
21
+ import json
22
+ import warnings
23
+ from argparse import ArgumentParser, Namespace
24
+ from importlib import import_module
25
+
26
+ import huggingface_hub
27
+ import torch
28
+ from huggingface_hub import hf_hub_download
29
+ from packaging import version
30
+
31
+ from ..utils import logging
32
+ from . import BaseDiffusersCLICommand
33
+
34
+
35
+ def conversion_command_factory(args: Namespace):
36
+ warnings.warn(
37
+ "`diffusers-cli fp16_safetensors` is deprecated and will be removed in a future version. "
38
+ "Convert weights to fp16 safetensors directly with `safetensors.torch.save_file` or via "
39
+ "`pipeline.save_pretrained(..., safe_serialization=True, variant='fp16')`.",
40
+ FutureWarning,
41
+ stacklevel=2,
42
+ )
43
+ if args.use_auth_token:
44
+ warnings.warn(
45
+ "The `--use_auth_token` flag is deprecated and will be removed in a future version."
46
+ "Authentication is now handled automatically if the user is logged in."
47
+ )
48
+ return FP16SafetensorsCommand(args.ckpt_id, args.fp16, args.use_safetensors)
49
+
50
+
51
+ class FP16SafetensorsCommand(BaseDiffusersCLICommand):
52
+ @staticmethod
53
+ def register_subcommand(parser: ArgumentParser):
54
+ conversion_parser = parser.add_parser(
55
+ "fp16_safetensors",
56
+ help="[DEPRECATED] Convert a Hub checkpoint's weights to fp16 safetensors and push back as a PR.",
57
+ usage="\n diffusers-cli fp16_safetensors [options]",
58
+ )
59
+ conversion_parser._optionals.title = "Options"
60
+ conversion_parser.add_argument(
61
+ "--ckpt_id",
62
+ type=str,
63
+ help="Repo id of the checkpoints on which to run the conversion. Example: 'openai/shap-e'.",
64
+ )
65
+ conversion_parser.add_argument(
66
+ "--fp16", action="store_true", help="If serializing the variables in FP16 precision."
67
+ )
68
+ conversion_parser.add_argument(
69
+ "--use_safetensors", action="store_true", help="If serializing in the safetensors format."
70
+ )
71
+ conversion_parser.add_argument(
72
+ "--use_auth_token",
73
+ action="store_true",
74
+ help="When working with checkpoints having private visibility. When used `hf auth login` needs to be run beforehand.",
75
+ )
76
+ conversion_parser.set_defaults(func=conversion_command_factory)
77
+
78
+ def __init__(self, ckpt_id: str, fp16: bool, use_safetensors: bool):
79
+ self.logger = logging.get_logger("diffusers-cli/fp16_safetensors")
80
+ self.ckpt_id = ckpt_id
81
+ self.local_ckpt_dir = f"/tmp/{ckpt_id}"
82
+ self.fp16 = fp16
83
+
84
+ self.use_safetensors = use_safetensors
85
+
86
+ if not self.use_safetensors and not self.fp16:
87
+ raise NotImplementedError(
88
+ "When `use_safetensors` and `fp16` both are False, then this command is of no use."
89
+ )
90
+
91
+ def run(self):
92
+ if version.parse(huggingface_hub.__version__) < version.parse("0.9.0"):
93
+ raise ImportError(
94
+ "The huggingface_hub version must be >= 0.9.0 to use this command. Please update your huggingface_hub"
95
+ " installation."
96
+ )
97
+ else:
98
+ from huggingface_hub import create_commit
99
+ from huggingface_hub._commit_api import CommitOperationAdd
100
+
101
+ model_index = hf_hub_download(repo_id=self.ckpt_id, filename="model_index.json")
102
+ with open(model_index, "r") as f:
103
+ pipeline_class_name = json.load(f)["_class_name"]
104
+ pipeline_class = getattr(import_module("diffusers"), pipeline_class_name)
105
+ self.logger.info(f"Pipeline class imported: {pipeline_class_name}.")
106
+
107
+ # Load the appropriate pipeline. We could have used `DiffusionPipeline`
108
+ # here, but just to avoid potential edge cases.
109
+ pipeline = pipeline_class.from_pretrained(
110
+ self.ckpt_id, torch_dtype=torch.float16 if self.fp16 else torch.float32
111
+ )
112
+ pipeline.save_pretrained(
113
+ self.local_ckpt_dir,
114
+ safe_serialization=True if self.use_safetensors else False,
115
+ variant="fp16" if self.fp16 else None,
116
+ )
117
+ self.logger.info(f"Pipeline locally saved to {self.local_ckpt_dir}.")
118
+
119
+ # Fetch all the paths.
120
+ if self.fp16:
121
+ modified_paths = glob.glob(f"{self.local_ckpt_dir}/*/*.fp16.*")
122
+ elif self.use_safetensors:
123
+ modified_paths = glob.glob(f"{self.local_ckpt_dir}/*/*.safetensors")
124
+
125
+ # Prepare for the PR.
126
+ commit_message = f"Serialize variables with FP16: {self.fp16} and safetensors: {self.use_safetensors}."
127
+ operations = []
128
+ for path in modified_paths:
129
+ operations.append(CommitOperationAdd(path_in_repo="/".join(path.split("/")[4:]), path_or_fileobj=path))
130
+
131
+ # Open the PR.
132
+ commit_description = (
133
+ "Variables converted by the [`diffusers`' `fp16_safetensors`"
134
+ " CLI](https://github.com/huggingface/diffusers/blob/main/src/diffusers/commands/fp16_safetensors.py)."
135
+ )
136
+ hub_pr_url = create_commit(
137
+ repo_id=self.ckpt_id,
138
+ operations=operations,
139
+ commit_message=commit_message,
140
+ commit_description=commit_description,
141
+ repo_type="model",
142
+ create_pr=True,
143
+ ).pr_url
144
+ self.logger.info(f"PR created here: {hub_pr_url}.")
diffusers/commands/run.py ADDED
@@ -0,0 +1,1227 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 The HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ """`diffusers-cli run` — single agentic entry point.
16
+
17
+ Runs any diffusers pipeline (standard or modular) by forwarding `--pipeline-kwargs` verbatim, saves the output by
18
+ detecting its runtime type, and can submit the same call to an HF Sandbox via `--remote`.
19
+ """
20
+
21
+ from __future__ import annotations
22
+
23
+ import json
24
+ import os
25
+ import sys
26
+ from argparse import ArgumentParser, Namespace, _SubParsersAction
27
+ from pathlib import Path
28
+ from typing import Any
29
+
30
+ from huggingface_hub.cli._output import out
31
+
32
+ from diffusers.models.attention_dispatch import _HUB_KERNELS_REGISTRY
33
+ from diffusers.utils import load_image, load_video, logging
34
+
35
+ from . import BaseDiffusersCLICommand
36
+
37
+
38
+ logger = logging.get_logger("diffusers-cli/run")
39
+
40
+
41
+ # ---------------------------------------------------------------------------
42
+ # Constants
43
+ # ---------------------------------------------------------------------------
44
+
45
+ DEFAULT_OUTPUT_DIR = str(Path.home() / ".diffusers" / "cli" / "run" / "outputs")
46
+ DTYPE_CHOICES = ("auto", "float16", "fp16", "bfloat16", "bf16", "float32", "fp32")
47
+ CPU_OFFLOAD_CHOICES = ("model", "group")
48
+
49
+
50
+ ATTENTION_BACKEND_CHOICES = ("default", *sorted(b.value for b in _HUB_KERNELS_REGISTRY))
51
+
52
+ # Kwarg keys whose string value gets auto-loaded before being passed to the pipeline call.
53
+ # Images resolve via `diffusers.utils.load_image` → PIL.Image.Image; videos resolve via
54
+ # `diffusers.utils.load_video` → list[PIL.Image.Image].
55
+ _IMAGE_INPUT_KEYS = (
56
+ "image",
57
+ "mask_image",
58
+ "control_image",
59
+ "ip_adapter_image",
60
+ "image_2",
61
+ )
62
+ _VIDEO_INPUT_KEYS = (
63
+ "video",
64
+ "control_video",
65
+ )
66
+ _AUDIO_INPUT_KEYS = (
67
+ "initial_audio_waveforms",
68
+ "reference_audio",
69
+ "src_audio",
70
+ )
71
+
72
+ # Pipeline attribute prefixes that identify a denoiser submodule. Matches base names
73
+ # (`transformer`, `unet`) and their numbered variants (`transformer_2`, etc.).
74
+ _DENOISER_COMPONENT_KEYS = ("transformer", "unet")
75
+
76
+ _DEFAULT_REMOTE_DEPS = (
77
+ "diffusers",
78
+ "accelerate",
79
+ "transformers",
80
+ "safetensors",
81
+ "sentencepiece", # required by several text-encoder tokenizers (T5, LLaMA, …)
82
+ "ftfy", # required by older CLIP text-encoder paths
83
+ )
84
+
85
+ # Base sandbox image — provides torch + CUDA so `uv pip install --system`
86
+ # only has to add the small Python deps. cuda12.8 is the highest cuda12.x tag
87
+ # below the HF Jobs host driver's CUDA 12.9 max.
88
+ _DEFAULT_REMOTE_IMAGE = "pytorch/pytorch:2.10.0-cuda12.8-cudnn9-runtime"
89
+
90
+ # Installed console-script name invoked inside the sandbox after the deps land.
91
+ _CONTAINER_CLI_BINARY = "diffusers-cli"
92
+
93
+ # Working directories inside the sandbox: local media from `--pipeline-kwargs` is uploaded
94
+ # under _SANDBOX_INPUTS_DIR, and the sandbox CLI is told to write its outputs under
95
+ # _SANDBOX_OUTPUTS_DIR so we can download them back afterwards.
96
+ _SANDBOX_INPUTS_DIR = "/tmp/diffusers-cli/inputs"
97
+ _SANDBOX_OUTPUTS_DIR = "/tmp/diffusers-cli/outputs"
98
+
99
+ RUN_ID_ENV = "DIFFUSERS_CLI_RUN_ID"
100
+
101
+ # Namespace keys that control *how* a remote run is dispatched, not what the sandbox CLI
102
+ # runs. They are stripped when forwarding argv to the sandbox.
103
+ REMOTE_KEYS = frozenset(
104
+ {
105
+ "remote",
106
+ "flavor",
107
+ "timeout",
108
+ "dependencies",
109
+ "namespace",
110
+ "image",
111
+ "keep_alive",
112
+ "sandbox_id",
113
+ "idle_timeout",
114
+ "volume",
115
+ "func",
116
+ "format", # top-level --format is a local rendering flag; never forward to the sandbox
117
+ }
118
+ )
119
+
120
+
121
+ # ---------------------------------------------------------------------------
122
+ # Argparse helpers
123
+ # ---------------------------------------------------------------------------
124
+
125
+
126
+ def _add_loading_arguments(parser: ArgumentParser) -> None:
127
+ parser.add_argument("--model", "-m", required=True, help="Model id on the Hugging Face Hub or local path.")
128
+ parser.add_argument(
129
+ "--device-map",
130
+ default=None,
131
+ help=(
132
+ "Component placement. Accepts a torch device string (`cuda`, `cuda:0`, `cpu`, `mps`), "
133
+ "`balanced` for pipeline-level auto-split across visible GPUs, or a JSON dict of "
134
+ '`{"<component>": <device>}` for explicit per-component placement. Auto-detected if omitted.'
135
+ ),
136
+ )
137
+ parser.add_argument("--dtype", default="auto", choices=DTYPE_CHOICES, help="Torch dtype for pipeline weights.")
138
+ parser.add_argument("--variant", default=None, help='Optional weight variant (e.g. "fp16").')
139
+ parser.add_argument("--revision", default=None, help="Model revision (branch, tag, or commit SHA).")
140
+ parser.add_argument("--token", default=None, help="Hugging Face token for gated/private models.")
141
+ parser.add_argument("--trust-remote-code", action="store_true", help="Allow custom code from the Hub.")
142
+ parser.add_argument(
143
+ "--lora",
144
+ action="append",
145
+ default=None,
146
+ metavar="JSON",
147
+ help=(
148
+ "JSON dict describing a LoRA adapter to attach after the pipeline loads. Repeat to stack "
149
+ 'multiple adapters. Format: \'{"lora_id": "<id>", "lora_scale": <float>}\'. `lora_scale` '
150
+ "defaults to 1.0; `adapter_name` is optional (auto-generated as `lora_<i>` when stacking)."
151
+ ),
152
+ )
153
+
154
+
155
+ def _add_optimization_arguments(parser: ArgumentParser) -> None:
156
+ parser.add_argument(
157
+ "--cpu-offload",
158
+ choices=CPU_OFFLOAD_CHOICES,
159
+ default=None,
160
+ help=(
161
+ "Offload pipeline components to CPU during inference. "
162
+ "'model' uses enable_model_cpu_offload, "
163
+ "'group' uses pipeline.enable_group_offload(leaf_level, use_stream=True)."
164
+ ),
165
+ )
166
+ parser.add_argument(
167
+ "--attention-backend",
168
+ choices=ATTENTION_BACKEND_CHOICES,
169
+ default="default",
170
+ help=(
171
+ "Override the attention backend on the transformer/UNet. "
172
+ "Only Hub-hosted kernels are exposed — they auto-download on first use."
173
+ ),
174
+ )
175
+ parser.add_argument("--vae-tiling", action="store_true", help="Enable VAE tiling (lower peak VRAM).")
176
+ parser.add_argument("--vae-slicing", action="store_true", help="Enable VAE slicing (lower peak VRAM).")
177
+ parser.add_argument(
178
+ "--context-parallel",
179
+ action="store_true",
180
+ help=(
181
+ "Enable Ulysses-style context parallelism (ulysses_anything mode). "
182
+ "Requires a DiT-based pipeline and launching the CLI under torchrun with ≥2 GPUs."
183
+ ),
184
+ )
185
+ parser.add_argument(
186
+ "--compile",
187
+ nargs="?",
188
+ const='{"fullgraph": true}',
189
+ default=None,
190
+ metavar="JSON",
191
+ help=(
192
+ "torch.compile every denoiser submodule on the pipeline. Accepts an optional JSON "
193
+ 'object of kwargs forwarded to `torch.compile`, e.g. \'{"mode": "max-autotune", '
194
+ '"fullgraph": true}\'. Bare `--compile` uses `fullgraph=true`. Adds a one-time '
195
+ "compilation cost on the first step but speeds up every subsequent step — worth it "
196
+ "for multi-step generation (50+ steps)."
197
+ ),
198
+ )
199
+
200
+
201
+ def _add_output_arguments(parser: ArgumentParser) -> None:
202
+ parser.add_argument(
203
+ "--output",
204
+ "-o",
205
+ default=None,
206
+ help=(
207
+ "Output file or directory. Defaults to "
208
+ "~/.diffusers/cli/run/outputs/diffusers-run-<YYYYMMDDTHHMMSS>-<short-uuid>/<NNNN>.<ext>."
209
+ ),
210
+ )
211
+ parser.add_argument(
212
+ "--push-to",
213
+ default=None,
214
+ help=(
215
+ "Upload the generated files to this HF bucket after saving (created if missing). Accepts "
216
+ "an HF bucket id (`<namespace>/<name>`), an `hf://buckets/<namespace>/<name>[/<subpath>]` "
217
+ "URI, or a browser URL for the same — a subpath is used as a folder prefix. Under --remote "
218
+ "the upload runs inside the sandbox; without an explicit --output the bucket becomes the "
219
+ "sole destination and nothing is downloaded back."
220
+ ),
221
+ )
222
+
223
+
224
+ def _add_remote_arguments(parser: ArgumentParser) -> None:
225
+ parser.add_argument(
226
+ "--remote",
227
+ action="store_true",
228
+ help="Run this command in a Hugging Face Sandbox instead of on the local machine.",
229
+ )
230
+ parser.add_argument(
231
+ "--flavor",
232
+ default="a10g-small",
233
+ help="HF Sandbox hardware flavor for --remote (e.g. a10g-small, a100-large, cpu-basic).",
234
+ )
235
+ parser.add_argument(
236
+ "--timeout",
237
+ default="10m",
238
+ help="Max wallclock for the run command inside the sandbox (e.g. 30m, 2h). Defaults to 10m.",
239
+ )
240
+ parser.add_argument(
241
+ "--dependencies",
242
+ action="append",
243
+ default=None,
244
+ help="Extra pip dependencies to install in the sandbox. Repeat to add multiple.",
245
+ )
246
+ parser.add_argument(
247
+ "--namespace",
248
+ default=None,
249
+ help="HF namespace to create the sandbox under (defaults to the current user).",
250
+ )
251
+ parser.add_argument(
252
+ "--image",
253
+ default=None,
254
+ help=(
255
+ "Sandbox image for --remote (defaults to "
256
+ f"{_DEFAULT_REMOTE_IMAGE!r}). Must provide torch + CUDA; the CLI installs the "
257
+ "small Python deps on top via `uv pip install --system`."
258
+ ),
259
+ )
260
+ parser.add_argument(
261
+ "--keep-alive",
262
+ action="store_true",
263
+ help=(
264
+ "Don't terminate the sandbox after the run. Its id is printed so a later --remote run "
265
+ "can reconnect with --sandbox-id and reuse the warm deps/weights/compile cache."
266
+ ),
267
+ )
268
+ parser.add_argument(
269
+ "--sandbox-id",
270
+ default=None,
271
+ help=(
272
+ "Reconnect to an existing sandbox (from a prior --keep-alive run) instead of creating a new "
273
+ "one, reusing its warm deps/weights/compile cache. Implies --keep-alive; stop it with "
274
+ "`hf sandbox kill <id>`."
275
+ ),
276
+ )
277
+ parser.add_argument(
278
+ "--idle-timeout",
279
+ default="10m",
280
+ help=(
281
+ "Auto-shutdown the sandbox after this much inactivity (e.g. 30m, 1h). Defaults to 10m. "
282
+ "Only applied on new sandbox creation — ignored when reconnecting via --sandbox-id."
283
+ ),
284
+ )
285
+ parser.add_argument(
286
+ "--volume",
287
+ action="append",
288
+ default=None,
289
+ metavar="BUCKET_ID[:MOUNT_PATH]",
290
+ help=(
291
+ "Mount an HF bucket into the sandbox as a read-write directory. Repeatable. Format: "
292
+ "`<namespace>/<name>` (mounts at `/mnt/buckets/<namespace>/<name>`) or "
293
+ "`<namespace>/<name>:/some/path` for a custom path. Reference mounted files from "
294
+ "--pipeline-kwargs like any other local path. Applied only on new sandbox creation — "
295
+ "ignored when reconnecting via --sandbox-id."
296
+ ),
297
+ )
298
+
299
+
300
+ # ---------------------------------------------------------------------------
301
+ # Pipeline loading + optimization
302
+ # ---------------------------------------------------------------------------
303
+
304
+
305
+ def _resolve_dtype(name: str | None):
306
+ if name in (None, "auto"):
307
+ return "auto"
308
+ import torch
309
+
310
+ mapping = {
311
+ "fp32": torch.float32,
312
+ "float32": torch.float32,
313
+ "fp16": torch.float16,
314
+ "float16": torch.float16,
315
+ "bf16": torch.bfloat16,
316
+ "bfloat16": torch.bfloat16,
317
+ }
318
+ if name not in mapping:
319
+ raise ValueError(f"Unknown dtype: {name}")
320
+ return mapping[name]
321
+
322
+
323
+ def _resolve_device_map(raw: str | None) -> str | dict:
324
+ """Parse `--device-map` into a value acceptable by `from_pretrained(device_map=...)`.
325
+
326
+ Returns a JSON dict if the value looks like one, `"balanced"` verbatim, or a single-device string (e.g. `"cuda"`,
327
+ `"cuda:1"`, `"cpu"`, `"mps"`). Auto-detects when `raw is None`, pinning to `cuda:$LOCAL_RANK` under torchrun.
328
+ """
329
+ if raw is None:
330
+ from diffusers.utils.torch_utils import torch_device
331
+
332
+ if torch_device == "cuda":
333
+ local_rank = os.environ.get("LOCAL_RANK")
334
+ if local_rank is not None:
335
+ import torch
336
+
337
+ torch.cuda.set_device(int(local_rank))
338
+ return f"cuda:{local_rank}"
339
+ return torch_device
340
+
341
+ if raw.strip().startswith("{"):
342
+ try:
343
+ parsed = json.loads(raw)
344
+ except json.JSONDecodeError as e:
345
+ raise SystemExit(f"--device-map must be a device string or a JSON dict: {e}") from e
346
+ if not isinstance(parsed, dict):
347
+ raise SystemExit("--device-map JSON must decode to an object.")
348
+ return parsed
349
+
350
+ return raw
351
+
352
+
353
+ def _apply_cpu_offload(pipeline: Any, mode: str, device_map: str | dict) -> None:
354
+ """Apply model or group CPU offload. Requires a single-device target (not balanced or dict)."""
355
+ if not isinstance(device_map, str) or device_map == "balanced":
356
+ raise SystemExit(
357
+ "--cpu-offload requires --device-map to be a single device string (e.g. 'cuda'); "
358
+ f"got {device_map!r}. balanced/dict placement is incompatible with CPU offload."
359
+ )
360
+
361
+ if mode == "model":
362
+ pipeline.enable_model_cpu_offload(device=device_map)
363
+ elif mode == "group":
364
+ import torch
365
+
366
+ pipeline.enable_group_offload(
367
+ onload_device=torch.device(device_map),
368
+ offload_type="leaf_level",
369
+ use_stream=True,
370
+ )
371
+
372
+
373
+ def _set_attention_backend(pipeline: Any, backend: str) -> None:
374
+ transformer = getattr(pipeline, "transformer", None)
375
+ if transformer is None or not hasattr(transformer, "set_attention_backend"):
376
+ logger.warning(
377
+ f"--attention-backend is only supported on transformer-based pipelines; "
378
+ f"{type(pipeline).__name__} uses the legacy UNet attention path."
379
+ )
380
+ return
381
+ try:
382
+ transformer.set_attention_backend(backend)
383
+ except (ValueError, ImportError, RuntimeError) as e:
384
+ logger.warning(
385
+ f"Attention backend {backend!r} could not be set on {type(transformer).__name__}: "
386
+ f"{type(e).__name__}: {e}. Falling back to the model's default backend."
387
+ )
388
+
389
+
390
+ def _enable_context_parallel(pipeline: Any) -> None:
391
+ import torch
392
+
393
+ if not torch.distributed.is_available():
394
+ raise SystemExit("--context-parallel requires a torch build with distributed support.")
395
+
396
+ if not torch.distributed.is_initialized():
397
+ # Hybrid backend: ulysses_anything's per-rank size coordination wants Gloo on CPU
398
+ # (avoids H2D/D2H for a tiny int tensor); the main attention all-to-all stays on NCCL.
399
+ torch.distributed.init_process_group(backend="cpu:gloo,cuda:nccl")
400
+
401
+ transformer = getattr(pipeline, "transformer", None)
402
+ if transformer is None or not hasattr(transformer, "enable_parallelism"):
403
+ raise SystemExit(
404
+ "--context-parallel requires a DiT-based pipeline. "
405
+ f"{type(pipeline).__name__} does not expose a `transformer` with `enable_parallelism`."
406
+ )
407
+
408
+ from diffusers import ContextParallelConfig
409
+
410
+ transformer.enable_parallelism(
411
+ config=ContextParallelConfig(
412
+ ulysses_degree=torch.distributed.get_world_size(),
413
+ ring_degree=1,
414
+ ulysses_anything=True,
415
+ )
416
+ )
417
+
418
+
419
+ def _apply_optimizations(pipeline: Any, args: Namespace) -> None:
420
+ """Apply VAE tiling/slicing, attention backend, context-parallel, and torch.compile toggles."""
421
+ vae = getattr(pipeline, "vae", None)
422
+ if args.vae_tiling and vae is not None and hasattr(vae, "enable_tiling"):
423
+ vae.enable_tiling()
424
+ if args.vae_slicing and vae is not None and hasattr(vae, "enable_slicing"):
425
+ vae.enable_slicing()
426
+ if args.attention_backend != "default":
427
+ _set_attention_backend(pipeline, args.attention_backend)
428
+ if args.context_parallel:
429
+ _enable_context_parallel(pipeline)
430
+ if args.compile is not None:
431
+ if args.context_parallel:
432
+ logger.warning("--compile is currently not supported with --context-parallel; skipping compile.")
433
+ else:
434
+ _compile_denoiser(pipeline, args.compile)
435
+
436
+
437
+ def _compile_denoiser(pipeline: Any, compile_spec: str) -> None:
438
+ """Compile every `transformer*` and `unet*` submodule on the pipeline.
439
+
440
+ `compile_spec` is the raw JSON string from `--compile` (`"{}"` for bare flag). Decoded into kwargs and forwarded
441
+ verbatim to the compile call.
442
+
443
+ Prefers regional compilation via `module.compile_repeated_blocks(**kwargs)` — only compiles the repeated inner
444
+ blocks (the bulk of the compute), much faster first-step latency than compiling the whole module. Falls back to
445
+ full `torch.compile` if the model doesn't expose `_repeated_blocks`.
446
+ """
447
+ import torch
448
+
449
+ try:
450
+ compile_kwargs = json.loads(compile_spec)
451
+ except json.JSONDecodeError as e:
452
+ raise SystemExit(f"--compile must be valid JSON: {e}") from e
453
+ if not isinstance(compile_kwargs, dict):
454
+ raise SystemExit("--compile must decode to a JSON object.")
455
+
456
+ for attr in dir(pipeline):
457
+ if not any(attr.startswith(key) for key in _DENOISER_COMPONENT_KEYS):
458
+ continue
459
+ module = getattr(pipeline, attr, None)
460
+ if not isinstance(module, torch.nn.Module):
461
+ continue
462
+
463
+ if getattr(module, "_repeated_blocks", None):
464
+ # Regional compile — only the repeated blocks. Mutates `module` in place.
465
+ module.compile_repeated_blocks(**compile_kwargs)
466
+ else:
467
+ # No regional metadata declared; fall back to compiling the whole module.
468
+ setattr(pipeline, attr, torch.compile(module, **compile_kwargs))
469
+
470
+
471
+ def _load_lora(pipeline: Any, args: Namespace) -> None:
472
+ """Attach one or more LoRA adapters. Each `--lora` value is a JSON dict.
473
+
474
+ Per-entry fields: `lora_id` (required), `lora_scale` (optional float, default 1.0), `adapter_name` (optional;
475
+ auto-generated as `lora_<i>` when stacking). Multiple `--lora` flags stack via a single `set_adapters(...)` call at
476
+ the end.
477
+ """
478
+ if not args.lora:
479
+ return
480
+ specs = []
481
+ for raw in args.lora:
482
+ try:
483
+ parsed = json.loads(raw)
484
+ except json.JSONDecodeError as e:
485
+ raise SystemExit(f"--lora must be valid JSON: {e}") from e
486
+ if not isinstance(parsed, dict):
487
+ raise SystemExit(f"--lora must decode to a JSON object; got {type(parsed).__name__}.")
488
+ specs.append(parsed)
489
+ if not hasattr(pipeline, "load_lora_weights"):
490
+ raise SystemExit(f"{type(pipeline).__name__} does not support LoRA loading.")
491
+
492
+ names: list[str] = []
493
+ scales: list[float] = []
494
+ for i, spec in enumerate(specs):
495
+ lora_id = spec.get("lora_id")
496
+ if not lora_id:
497
+ raise SystemExit(f"--lora entry {i} is missing 'lora_id'.")
498
+ adapter_name = spec.get("adapter_name") or (f"lora_{i}" if len(specs) > 1 else "default")
499
+ pipeline.load_lora_weights(lora_id, adapter_name=adapter_name)
500
+ names.append(adapter_name)
501
+ scales.append(float(spec.get("lora_scale", 1.0)))
502
+
503
+ if hasattr(pipeline, "set_adapters"):
504
+ pipeline.set_adapters(names, adapter_weights=scales)
505
+
506
+
507
+ def _load_pipeline(args: Namespace) -> Any:
508
+ import diffusers
509
+
510
+ # Detect modular repos by trying the standard config; `ModularPipeline` repos ship
511
+ # `modular_model_index.json` instead of `model_index.json`, so `load_config` OSErrors.
512
+ try:
513
+ diffusers.DiffusionPipeline.load_config(args.model, token=args.token, revision=args.revision)
514
+ modular = False
515
+ except OSError:
516
+ modular = True
517
+
518
+ dtype = _resolve_dtype(args.dtype)
519
+ device_map = _resolve_device_map(args.device_map)
520
+ common_kwargs: dict[str, Any] = {
521
+ "trust_remote_code": args.trust_remote_code,
522
+ }
523
+ if dtype != "auto":
524
+ common_kwargs["torch_dtype"] = dtype
525
+ if args.variant:
526
+ common_kwargs["variant"] = args.variant
527
+ if args.token:
528
+ common_kwargs["token"] = args.token
529
+ # CPU offload sets up its own placement hooks, so leave weights on CPU at load time.
530
+ if not args.cpu_offload:
531
+ common_kwargs["device_map"] = device_map
532
+
533
+ if modular:
534
+ # ModularPipeline.from_pretrained fetches only the pipeline config; component
535
+ # weights come in via load_components(). `revision` scopes the config fetch,
536
+ # so it stays on from_pretrained — each ComponentSpec pins its own revision,
537
+ # and forwarding a global `revision` to load_components() would override those.
538
+ pipeline = diffusers.ModularPipeline.from_pretrained(
539
+ args.model,
540
+ trust_remote_code=args.trust_remote_code,
541
+ token=args.token,
542
+ revision=args.revision,
543
+ )
544
+ pipeline.load_components(**common_kwargs)
545
+ else:
546
+ pipeline = diffusers.DiffusionPipeline.from_pretrained(args.model, revision=args.revision, **common_kwargs)
547
+
548
+ _load_lora(pipeline, args)
549
+ if args.cpu_offload:
550
+ _apply_cpu_offload(pipeline, args.cpu_offload, device_map)
551
+ _apply_optimizations(pipeline, args)
552
+
553
+ return pipeline
554
+
555
+
556
+ # ---------------------------------------------------------------------------
557
+ # Pipeline call helpers
558
+ # ---------------------------------------------------------------------------
559
+
560
+
561
+ def _parse_pipeline_kwargs(raw: str | None) -> dict[str, Any]:
562
+ if not raw:
563
+ return {}
564
+ try:
565
+ parsed = json.loads(raw)
566
+ except json.JSONDecodeError as e:
567
+ raise SystemExit(f"--pipeline-kwargs must be valid JSON: {e}") from e
568
+ if not isinstance(parsed, dict):
569
+ raise SystemExit("--pipeline-kwargs must decode to a JSON object.")
570
+ return parsed
571
+
572
+
573
+ def _load_audio(url_or_path: str) -> tuple[Any, int]:
574
+ """Load audio from a URL or local path via torchaudio. Returns `(waveform, sampling_rate)`."""
575
+ import torchaudio
576
+
577
+ if url_or_path.startswith(("http://", "https://")):
578
+ import io
579
+
580
+ import httpx
581
+
582
+ from ..utils.constants import DIFFUSERS_REQUEST_TIMEOUT
583
+
584
+ resp = httpx.get(url_or_path, follow_redirects=True, timeout=DIFFUSERS_REQUEST_TIMEOUT)
585
+ resp.raise_for_status()
586
+ return torchaudio.load(io.BytesIO(resp.content))
587
+ return torchaudio.load(url_or_path)
588
+
589
+
590
+ def _resolve_media_inputs(call_kwargs: dict[str, Any]) -> None:
591
+ """Replace string paths/URLs at known media-input keys with loaded tensors.
592
+
593
+ Images resolve to `PIL.Image.Image` via `load_image`; videos to `list[PIL.Image.Image]` via `load_video`; audio to
594
+ a `torch.Tensor` via `_load_audio` (also auto-sets the paired sampling-rate kwarg for `initial_audio_waveforms` if
595
+ the user didn't supply it). A `list[str]` at any key is treated as a batch: each entry is loaded and the value
596
+ becomes a list of loaded objects. Non-string, non-list values pass through untouched.
597
+ """
598
+
599
+ def _is_string_list(v: Any) -> bool:
600
+ return isinstance(v, list) and bool(v) and all(isinstance(x, str) for x in v)
601
+
602
+ for key in _IMAGE_INPUT_KEYS:
603
+ value = call_kwargs.get(key)
604
+ if isinstance(value, str):
605
+ call_kwargs[key] = load_image(value)
606
+ elif _is_string_list(value):
607
+ call_kwargs[key] = [load_image(v) for v in value]
608
+ for key in _VIDEO_INPUT_KEYS:
609
+ value = call_kwargs.get(key)
610
+ if isinstance(value, str):
611
+ call_kwargs[key] = load_video(value)
612
+ elif _is_string_list(value):
613
+ call_kwargs[key] = [load_video(v) for v in value]
614
+ for key in _AUDIO_INPUT_KEYS:
615
+ value = call_kwargs.get(key)
616
+ if isinstance(value, str):
617
+ waveform, sr = _load_audio(value)
618
+ call_kwargs[key] = waveform
619
+ if key == "initial_audio_waveforms" and "initial_audio_sampling_rate" not in call_kwargs:
620
+ call_kwargs["initial_audio_sampling_rate"] = sr
621
+ elif _is_string_list(value):
622
+ pairs = [_load_audio(v) for v in value]
623
+ call_kwargs[key] = [w for w, _ in pairs]
624
+ if key == "initial_audio_waveforms" and "initial_audio_sampling_rate" not in call_kwargs:
625
+ # All batched waveforms must share a sampling rate; use the first entry's.
626
+ call_kwargs["initial_audio_sampling_rate"] = pairs[0][1]
627
+
628
+
629
+ def _get_generator(seed: int | None, device: str):
630
+ if seed is None:
631
+ return None
632
+ import torch
633
+
634
+ generator_device = "cpu" if device == "mps" else device
635
+ return torch.Generator(device=generator_device).manual_seed(seed)
636
+
637
+
638
+ def _unwrap_pipeline_output(result: Any) -> Any:
639
+ """Unwrap a pipeline-output object into the raw payload the saver can dispatch on."""
640
+ if hasattr(result, "images"):
641
+ return result.images
642
+ if hasattr(result, "frames"):
643
+ return result.frames[0]
644
+ if hasattr(result, "audios"):
645
+ return result.audios
646
+ return result
647
+
648
+
649
+ # ---------------------------------------------------------------------------
650
+ # Output saving (dispatch by type)
651
+ # ---------------------------------------------------------------------------
652
+
653
+
654
+ def _get_or_create_run_id() -> str:
655
+ """Return the current run's id, creating one if not yet set.
656
+
657
+ Format: `diffusers-run-<YYYYMMDDTHHMMSS>-<6-char-uuid>`. Same id is reused as the local output subdirectory, the
658
+ remote bucket prefix, and the container-side `RUN_ID_ENV` so a run's artifacts are traceable end-to-end.
659
+ """
660
+ import uuid
661
+ from datetime import datetime
662
+
663
+ existing = os.environ.get(RUN_ID_ENV)
664
+ if existing:
665
+ return existing
666
+ run_id = f"diffusers-run-{datetime.now().strftime('%Y%m%dT%H%M%S')}-{uuid.uuid4().hex[:6]}"
667
+ os.environ[RUN_ID_ENV] = run_id
668
+ return run_id
669
+
670
+
671
+ def _resolve_output_paths(task: str, num: int, explicit: str | None, ext: str) -> list[Path]:
672
+ if explicit is None:
673
+ base = Path(DEFAULT_OUTPUT_DIR) / _get_or_create_run_id()
674
+ base.mkdir(parents=True, exist_ok=True)
675
+ return [base / f"{i:04d}.{ext}" for i in range(num)]
676
+
677
+ p = Path(explicit)
678
+ if explicit.endswith(os.sep) or p.is_dir():
679
+ p.mkdir(parents=True, exist_ok=True)
680
+ return [p / f"{i:04d}.{ext}" for i in range(num)]
681
+
682
+ p.parent.mkdir(parents=True, exist_ok=True)
683
+ if num == 1:
684
+ return [p]
685
+ stem, suffix = p.stem, p.suffix or f".{ext}"
686
+ return [p.with_name(f"{stem}-{i:04d}{suffix}") for i in range(num)]
687
+
688
+
689
+ def _as_pil_list(value: Any):
690
+ try:
691
+ from PIL.Image import Image as PILImage
692
+ except ImportError:
693
+ return None
694
+ if isinstance(value, PILImage):
695
+ return [value]
696
+ if isinstance(value, (list, tuple)) and value and all(isinstance(v, PILImage) for v in value):
697
+ return list(value)
698
+ return None
699
+
700
+
701
+ def _as_frame_sequence(value: Any):
702
+ try:
703
+ from PIL.Image import Image as PILImage
704
+ except ImportError:
705
+ PILImage = None # type: ignore[assignment]
706
+
707
+ if isinstance(value, (list, tuple)) and len(value) >= 2:
708
+ first = value[0]
709
+ if PILImage is not None and isinstance(first, PILImage):
710
+ return list(value)
711
+ try:
712
+ import numpy as np
713
+
714
+ if isinstance(first, np.ndarray):
715
+ return list(value)
716
+ except ImportError:
717
+ pass
718
+ return None
719
+
720
+
721
+ def _as_audio_arrays(value: Any):
722
+ try:
723
+ import numpy as np
724
+ except ImportError:
725
+ return None
726
+ if isinstance(value, np.ndarray) and value.ndim <= 2:
727
+ return [value]
728
+ if isinstance(value, (list, tuple)) and value and all(isinstance(v, np.ndarray) for v in value):
729
+ return list(value)
730
+ return None
731
+
732
+
733
+ def _save_audio_arrays(audios, sampling_rate: int, args: Namespace, task: str) -> list[str]:
734
+ """Write each numpy audio array to a 16-bit PCM WAV at `sampling_rate` Hz.
735
+
736
+ Uses the stdlib `wave` module so no scipy dependency is required.
737
+ """
738
+ import wave
739
+
740
+ import numpy as np
741
+
742
+ paths = _resolve_output_paths(task, len(audios), args.output, ext="wav")
743
+ saved: list[str] = []
744
+ for audio, path in zip(audios, paths):
745
+ data = np.asarray(audio)
746
+ if data.dtype.kind == "f":
747
+ data = (np.clip(data, -1.0, 1.0) * 32767).astype(np.int16)
748
+ else:
749
+ data = data.astype(np.int16)
750
+ if data.ndim == 1:
751
+ n_channels = 1
752
+ else:
753
+ # Heuristic: shorter axis is channels (interleaved layout for `wave` is
754
+ # samples × channels, so transpose if needed).
755
+ if data.shape[0] < data.shape[-1]:
756
+ data = data.T
757
+ n_channels = data.shape[1]
758
+ with wave.open(str(path), "wb") as w:
759
+ w.setnchannels(n_channels)
760
+ w.setsampwidth(2) # 16-bit PCM
761
+ w.setframerate(sampling_rate)
762
+ w.writeframes(data.tobytes())
763
+ saved.append(str(path))
764
+ return saved
765
+
766
+
767
+ def _save_output(value: Any, args: Namespace, task: str) -> list[str]:
768
+ """Save `value` by dispatching on its runtime type."""
769
+ pil_images = _as_pil_list(value)
770
+ if pil_images is not None:
771
+ paths = _resolve_output_paths(task, len(pil_images), args.output, ext="png")
772
+ for img, path in zip(pil_images, paths):
773
+ img.save(path)
774
+ return [str(p) for p in paths]
775
+
776
+ frames = _as_frame_sequence(value)
777
+ if frames is not None:
778
+ from diffusers.utils import export_to_video
779
+
780
+ path = _resolve_output_paths(task, 1, args.output, ext="mp4")[0]
781
+ export_to_video(frames, str(path), fps=args.fps)
782
+ return [str(path)]
783
+
784
+ audios = _as_audio_arrays(value)
785
+ if audios is not None:
786
+ return _save_audio_arrays(audios, args.sampling_rate or 16000, args, task)
787
+
788
+ path = _resolve_output_paths(task, 1, args.output, ext="json")[0]
789
+ Path(path).write_text(json.dumps(value, default=str, indent=2))
790
+ return [str(path)]
791
+
792
+
793
+ # ---------------------------------------------------------------------------
794
+ # Hub bucket upload (--push-to)
795
+ # ---------------------------------------------------------------------------
796
+
797
+
798
+ def _parse_push_to(spec: str) -> tuple[str, str]:
799
+ """Split `--push-to` into a bucket id and an optional subpath prefix.
800
+
801
+ Accepts an HF bucket id (`<namespace>/<name>[/<subpath>]`), a canonical
802
+ `hf://buckets/<namespace>/<name>[/<subpath>]` URI, or a Hub web URL for the same. Non-bucket URIs (models,
803
+ datasets, spaces) are rejected — `--push-to` targets storage buckets only.
804
+ """
805
+ from huggingface_hub import parse_hf_uri
806
+
807
+ # Bare shorthand → canonical URI so a single parser handles every accepted form.
808
+ if not spec.startswith(("hf://", "http://", "https://")):
809
+ spec = f"hf://buckets/{spec.strip('/')}"
810
+ uri = parse_hf_uri(spec)
811
+ if not uri.is_bucket:
812
+ raise SystemExit(f"--push-to must point at a bucket; got {uri.type!r} URI {spec!r}.")
813
+ return uri.id, uri.path_in_repo
814
+
815
+
816
+ def _push_outputs(args: Namespace, saved_paths: list[str], task: str) -> dict[str, Any] | None:
817
+ """Upload `saved_paths` to the `--push-to` bucket. Returns a summary or None."""
818
+ if not args.push_to:
819
+ return None
820
+
821
+ from huggingface_hub import HfApi
822
+
823
+ bucket_id, subpath = _parse_push_to(args.push_to)
824
+ api = HfApi(token=args.token)
825
+ api.create_bucket(bucket_id, exist_ok=True)
826
+
827
+ run_id = _get_or_create_run_id()
828
+ prefix = f"{subpath}/{run_id}" if subpath else run_id
829
+ add = [(local, f"{prefix}/{Path(local).name}") for local in saved_paths]
830
+ api.batch_bucket_files(bucket_id, add=add)
831
+
832
+ uploaded = [f"hf://buckets/{bucket_id}/{dest}" for _, dest in add]
833
+ return {"bucket_id": bucket_id, "uploaded": uploaded}
834
+
835
+
836
+ # ---------------------------------------------------------------------------
837
+ # Remote execution (HF Sandbox)
838
+ # ---------------------------------------------------------------------------
839
+
840
+
841
+ def _build_task_kwargs(args: Namespace) -> dict[str, Any]:
842
+ """Pick out the kwargs the sandbox CLI should invoke the task with."""
843
+ out: dict[str, Any] = {}
844
+ for key, value in vars(args).items():
845
+ if key in REMOTE_KEYS or value is None or value is False:
846
+ continue
847
+ out[key] = value
848
+ return out
849
+
850
+
851
+ def _kwargs_to_argv(task: str, task_kwargs: dict[str, Any]) -> list[str]:
852
+ """Render `task_kwargs` as the argv list the sandbox CLI's argparse will see."""
853
+ argv: list[str] = [task]
854
+ for key, value in task_kwargs.items():
855
+ flag = "--" + key.replace("_", "-")
856
+ if value is True:
857
+ argv.append(flag)
858
+ elif isinstance(value, list):
859
+ for item in value:
860
+ argv.extend([flag, str(item)])
861
+ else:
862
+ argv.extend([flag, str(value)])
863
+ return argv
864
+
865
+
866
+ def _duration_to_seconds(value: str) -> float:
867
+ """Parse a duration like `30s`, `10m`, `2h` (or a bare number of seconds) into seconds."""
868
+ value = value.strip()
869
+ units = {"s": 1, "m": 60, "h": 3600}
870
+ if value and value[-1] in units:
871
+ return float(value[:-1]) * units[value[-1]]
872
+ return float(value)
873
+
874
+
875
+ def _upload_inputs_to_sandbox(args: Namespace, sbx: Any, run_id: str) -> None:
876
+ """Upload local media paths in `--pipeline-kwargs` into the sandbox and rewrite the JSON in place.
877
+
878
+ Walks known image/video/audio-input keys; any string value that resolves to a local file is uploaded to
879
+ `<_SANDBOX_INPUTS_DIR>/<run_id>/<key>_<basename>` and the JSON path is rewritten to that in-sandbox path. URLs,
880
+ `hf://` URIs, and non-existent paths pass through untouched.
881
+ """
882
+ if not args.pipeline_kwargs:
883
+ return
884
+ try:
885
+ parsed = json.loads(args.pipeline_kwargs)
886
+ except json.JSONDecodeError:
887
+ return # the sandbox CLI will fail loudly with a parse error later
888
+ if not isinstance(parsed, dict):
889
+ return
890
+
891
+ def _upload_one(key: str, index: int | None, local_str: str) -> str:
892
+ # `index` is None for scalar entries, an int for list entries (used to disambiguate names).
893
+ local = Path(local_str)
894
+ suffix = f"_{index}" if index is not None else ""
895
+ remote_path = f"{_SANDBOX_INPUTS_DIR}/{run_id}/{key}{suffix}_{local.name}"
896
+ sbx.files.upload(str(local), remote_path)
897
+ return remote_path
898
+
899
+ uploaded = 0
900
+ for key in (*_IMAGE_INPUT_KEYS, *_VIDEO_INPUT_KEYS, *_AUDIO_INPUT_KEYS):
901
+ value = parsed.get(key)
902
+ if isinstance(value, str) and Path(value).is_file():
903
+ parsed[key] = _upload_one(key, None, value)
904
+ uploaded += 1
905
+ elif isinstance(value, list):
906
+ # Batched inputs: upload each local path, leave URLs/hf:// URIs alone.
907
+ new_list = list(value)
908
+ for i, entry in enumerate(value):
909
+ if isinstance(entry, str) and Path(entry).is_file():
910
+ new_list[i] = _upload_one(key, i, entry)
911
+ uploaded += 1
912
+ parsed[key] = new_list
913
+
914
+ if uploaded:
915
+ logger.info(f"uploaded {uploaded} local input file(s) to the sandbox")
916
+ args.pipeline_kwargs = json.dumps(parsed)
917
+
918
+
919
+ def _download_outputs_from_sandbox(sbx: Any, sandbox_dir: str, local_dir: Path) -> list[str]:
920
+ """Download every file the sandbox CLI wrote under `sandbox_dir` into `local_dir`."""
921
+ local_dir.mkdir(parents=True, exist_ok=True)
922
+ saved: list[str] = []
923
+ for entry in sbx.files.list(sandbox_dir):
924
+ if entry.type != "file":
925
+ continue
926
+ target = local_dir / Path(entry.path).name
927
+ sbx.files.download(entry.path, str(target))
928
+ saved.append(str(target))
929
+ return saved
930
+
931
+
932
+ def _maybe_submit_remote(args: Namespace, task: str) -> bool:
933
+ """If `--remote` was set, run this invocation inside an HF Sandbox and return True."""
934
+ if not args.remote:
935
+ return False
936
+
937
+ import shlex
938
+ import time
939
+
940
+ from huggingface_hub import get_token
941
+ from huggingface_hub.utils import send_telemetry
942
+
943
+ import diffusers
944
+
945
+ try:
946
+ from huggingface_hub import Sandbox
947
+ except ImportError:
948
+ raise SystemExit(
949
+ "--remote requires huggingface_hub>=1.23 for HF Sandbox support. "
950
+ "Upgrade with `pip install -U huggingface_hub`."
951
+ )
952
+
953
+ if Path(args.model).exists():
954
+ raise SystemExit(
955
+ f"--model {args.model!r} is a local path; the sandbox can't see it. "
956
+ "Pass a Hub repo id so the sandbox can download it."
957
+ )
958
+
959
+ hf_token = args.token or get_token()
960
+ run_id = _get_or_create_run_id()
961
+
962
+ # An explicit --push-to means the bucket is the user's destination, so skip the local
963
+ # download unless they also asked for a local path via --output.
964
+ user_bucket = bool(args.push_to)
965
+ download_locally = (not user_bucket) or (args.output is not None)
966
+ local_dir = Path(args.output) if args.output else Path(DEFAULT_OUTPUT_DIR) / run_id
967
+
968
+ use_existing_sandbox = bool(args.sandbox_id)
969
+ keep_alive = args.keep_alive or use_existing_sandbox
970
+ if use_existing_sandbox and args.volume:
971
+ logger.warning(
972
+ "--volume is ignored when reconnecting to an existing sandbox (mounts are set at creation time)."
973
+ )
974
+ if use_existing_sandbox:
975
+ logger.info(f"reconnecting to sandbox {args.sandbox_id!r}...")
976
+ sbx = Sandbox.connect(args.sandbox_id, token=hf_token)
977
+ else:
978
+ logger.info(f"creating sandbox on flavor={args.flavor!r}...")
979
+ create_kwargs: dict[str, Any] = {
980
+ "image": args.image or _DEFAULT_REMOTE_IMAGE,
981
+ "flavor": args.flavor,
982
+ "forward_hf_token": True,
983
+ "token": hf_token,
984
+ "env": {
985
+ "HF_ENABLE_PARALLEL_LOADING": "1",
986
+ "DIFFUSERS_VERBOSITY": os.environ.get("DIFFUSERS_VERBOSITY", "info"),
987
+ },
988
+ "idle_timeout": args.idle_timeout,
989
+ }
990
+ if args.volume:
991
+ from huggingface_hub import Volume
992
+
993
+ volumes = []
994
+ for spec in args.volume:
995
+ bucket_id, sep, mount_path = spec.partition(":")
996
+ if not sep:
997
+ mount_path = f"/mnt/buckets/{bucket_id}"
998
+ if bucket_id.count("/") != 1:
999
+ raise SystemExit(f"--volume: bucket id must be <namespace>/<name>, got {bucket_id!r}")
1000
+ if not mount_path.startswith("/"):
1001
+ raise SystemExit(f"--volume: mount path must be absolute, got {mount_path!r}")
1002
+ volumes.append(Volume(type="bucket", source=bucket_id, mount_path=mount_path))
1003
+ create_kwargs["volumes"] = volumes
1004
+ if args.namespace is not None:
1005
+ create_kwargs["namespace"] = args.namespace
1006
+ sbx = Sandbox.create(**create_kwargs)
1007
+
1008
+ def _stream(chunk: str) -> None:
1009
+ sys.stderr.write(chunk)
1010
+ sys.stderr.flush()
1011
+
1012
+ exit_code = 0
1013
+ saved: list[str] = []
1014
+ run_seconds = 0.0
1015
+ try:
1016
+ _upload_inputs_to_sandbox(args, sbx, run_id)
1017
+
1018
+ dependencies = list(_DEFAULT_REMOTE_DEPS)
1019
+ if args.dependencies:
1020
+ dependencies.extend(args.dependencies)
1021
+ # --break-system-packages bypasses PEP 668; harmless in a throwaway sandbox. uv is a
1022
+ # near no-op when the deps are already satisfied, so this stays cheap on a reused sandbox.
1023
+ install_cmd = shlex.join(["uv", "pip", "install", "--system", "--break-system-packages", *dependencies])
1024
+ logger.info("installing dependencies in the sandbox...")
1025
+ sbx.run(install_cmd, on_stdout=_stream, on_stderr=_stream)
1026
+
1027
+ # Per-run outputs subdirectory so a reused sandbox doesn't leak files from prior runs
1028
+ # into this run's download set.
1029
+ sandbox_output_dir = f"{_SANDBOX_OUTPUTS_DIR}/{run_id}"
1030
+ task_kwargs = _build_task_kwargs(args)
1031
+ task_kwargs["output"] = sandbox_output_dir + "/"
1032
+ cli_argv = _kwargs_to_argv(task, task_kwargs)
1033
+ # Suppress the container CLI's own `out.result(...)` payload — the outer wrapper owns the
1034
+ # final structured output for --remote runs.
1035
+ format_argv = ["--format", "quiet"]
1036
+ # torchrun wraps the CLI for --context-parallel so torch.distributed initializes across
1037
+ # every visible GPU before the run command starts.
1038
+ if args.context_parallel:
1039
+ cli_argv = [
1040
+ "torchrun",
1041
+ "--nproc-per-node=gpu",
1042
+ "-m",
1043
+ "diffusers.commands.diffusers_cli",
1044
+ *format_argv,
1045
+ *cli_argv,
1046
+ ]
1047
+ else:
1048
+ cli_argv = [_CONTAINER_CLI_BINARY, *format_argv, *cli_argv]
1049
+
1050
+ started = time.perf_counter()
1051
+ # Per-invocation env: RUN_ID_ENV must be fresh each run. Sandbox.create-time env is
1052
+ # baked in and would go stale on reused sandboxes, silently reusing the initial run's
1053
+ # bucket prefix in `_push_outputs`.
1054
+ result = sbx.run(
1055
+ cli_argv,
1056
+ env={RUN_ID_ENV: run_id},
1057
+ on_stdout=_stream,
1058
+ on_stderr=_stream,
1059
+ timeout=_duration_to_seconds(args.timeout),
1060
+ check=False,
1061
+ )
1062
+ run_seconds = time.perf_counter() - started
1063
+ exit_code = result.exit_code
1064
+
1065
+ if exit_code == 0 and download_locally:
1066
+ saved = _download_outputs_from_sandbox(sbx, sandbox_output_dir, local_dir)
1067
+ finally:
1068
+ if keep_alive:
1069
+ logger.info(
1070
+ f"sandbox {sbx.id} kept alive — reconnect with "
1071
+ f"`--remote --sandbox-id {sbx.id}`, stop with `hf sandbox kill {sbx.id}`."
1072
+ )
1073
+ else:
1074
+ sbx.kill()
1075
+
1076
+ send_telemetry(
1077
+ topic="diffusers/cli/run/remote",
1078
+ library_name="diffusers",
1079
+ library_version=diffusers.__version__,
1080
+ )
1081
+
1082
+ payload: dict[str, Any] = {
1083
+ "exit_code": exit_code,
1084
+ "run_seconds": round(run_seconds, 1),
1085
+ }
1086
+ if keep_alive:
1087
+ payload["sandbox_id"] = sbx.id
1088
+ if download_locally:
1089
+ payload["outputs"] = saved
1090
+ if args.push_to:
1091
+ bucket_id, subpath = _parse_push_to(args.push_to)
1092
+ prefix = f"{subpath}/{run_id}" if subpath else run_id
1093
+ payload["pushed-to"] = f"hf://buckets/{bucket_id}/{prefix}/"
1094
+ out.result("remote-run", **payload)
1095
+
1096
+ if exit_code != 0:
1097
+ raise SystemExit(f"remote run failed with exit code {exit_code}")
1098
+ return True
1099
+
1100
+
1101
+ # ---------------------------------------------------------------------------
1102
+ # Subcommand
1103
+ # ---------------------------------------------------------------------------
1104
+
1105
+
1106
+ class RunCommand(BaseDiffusersCLICommand):
1107
+ task = "run"
1108
+
1109
+ @staticmethod
1110
+ def register_subcommand(subparsers: _SubParsersAction) -> None:
1111
+ from argparse import RawDescriptionHelpFormatter
1112
+
1113
+ epilog = (
1114
+ "Examples\n"
1115
+ " $ diffusers-cli run -m black-forest-labs/FLUX.1-dev --dtype bf16 \\\n"
1116
+ ' --pipeline-kwargs \'{"prompt": "a cat on the moon"}\'\n'
1117
+ " $ diffusers-cli run -m black-forest-labs/FLUX.1-dev --dtype bf16 \\\n"
1118
+ ' --pipeline-kwargs \'{"prompt": "make the fur grey", "image": "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/cat.png"}\'\n'
1119
+ " $ diffusers-cli run -m black-forest-labs/FLUX.1-dev --dtype bf16 \\\n"
1120
+ ' --pipeline-kwargs \'{"prompt": "a tiny cat"}\' \\\n'
1121
+ ' --lora \'{"lora_id": "alvdansen/littletinies", "lora_scale": 0.8}\'\n'
1122
+ " $ diffusers-cli run -m black-forest-labs/FLUX.1-dev --dtype bf16 \\\n"
1123
+ ' --pipeline-kwargs \'{"prompt": "a cat"}\' --remote --flavor a100-large\n'
1124
+ " $ diffusers-cli run -m black-forest-labs/FLUX.1-dev --dtype bf16 --context-parallel \\\n"
1125
+ ' --pipeline-kwargs \'{"prompt": "a cat"}\' --remote --flavor 4xa100-large\n'
1126
+ "\n"
1127
+ "Learn more\n"
1128
+ " Use `diffusers-cli <command> --help` for more information about a command.\n"
1129
+ " Read the documentation at https://huggingface.co/docs/diffusers\n"
1130
+ )
1131
+
1132
+ parser: ArgumentParser = subparsers.add_parser(
1133
+ "run",
1134
+ help="Run any diffusers pipeline locally or remotely in an HF Sandbox.",
1135
+ usage="\n diffusers-cli run [options]",
1136
+ epilog=epilog,
1137
+ formatter_class=RawDescriptionHelpFormatter,
1138
+ )
1139
+ parser._optionals.title = "Options"
1140
+ _add_loading_arguments(parser)
1141
+ _add_optimization_arguments(parser)
1142
+ parser.add_argument(
1143
+ "--pipeline-kwargs",
1144
+ default=None,
1145
+ help=(
1146
+ "JSON object of kwargs passed to the pipeline call. String values at known "
1147
+ f"image-input keys ({', '.join(_IMAGE_INPUT_KEYS)}) are auto-loaded as PIL images; "
1148
+ f"video-input keys ({', '.join(_VIDEO_INPUT_KEYS)}) are auto-loaded as frame lists; "
1149
+ f"audio-input keys ({', '.join(_AUDIO_INPUT_KEYS)}) are auto-loaded via torchaudio."
1150
+ ),
1151
+ )
1152
+ parser.add_argument(
1153
+ "--output-key",
1154
+ default=None,
1155
+ help="For modular pipelines: name of the intermediate to extract (passed as `output=` to the call).",
1156
+ )
1157
+ parser.add_argument("--seed", type=int, default=None, help="Random seed for reproducibility.")
1158
+ parser.add_argument(
1159
+ "--fps",
1160
+ type=int,
1161
+ default=8,
1162
+ help="FPS used when the output happens to be a frame sequence.",
1163
+ )
1164
+ parser.add_argument(
1165
+ "--sampling-rate",
1166
+ type=int,
1167
+ default=None,
1168
+ help="Sample rate used when the output happens to be an audio array.",
1169
+ )
1170
+ _add_remote_arguments(parser)
1171
+ _add_output_arguments(parser)
1172
+ parser.set_defaults(func=RunCommand)
1173
+
1174
+ def __init__(self, args: Namespace):
1175
+ self.args = args
1176
+
1177
+ def run(self) -> None:
1178
+ import diffusers
1179
+
1180
+ _get_or_create_run_id() # populate RUN_ID_ENV so local output dir + remote bucket prefix agree
1181
+
1182
+ call_kwargs = _parse_pipeline_kwargs(self.args.pipeline_kwargs)
1183
+
1184
+ if _maybe_submit_remote(self.args, self.task):
1185
+ return
1186
+
1187
+ # Resolve media before loading pipeline weights so dead URLs / missing files fail
1188
+ # fast — cheap to fetch, expensive to load a 20GB model just to hit a 404.
1189
+ _resolve_media_inputs(call_kwargs)
1190
+ pipeline = _load_pipeline(self.args)
1191
+ is_modular = isinstance(pipeline, diffusers.ModularPipeline)
1192
+
1193
+ if self.args.output_key is not None:
1194
+ call_kwargs["output"] = self.args.output_key
1195
+
1196
+ device = pipeline.device.type if hasattr(pipeline, "device") else "cpu"
1197
+ generator = _get_generator(self.args.seed, device)
1198
+ if generator is not None:
1199
+ call_kwargs["generator"] = generator
1200
+
1201
+ try:
1202
+ result = pipeline(**call_kwargs)
1203
+
1204
+ # Under torchrun, ranks > 0 produce identical output to rank 0 (CP shards the
1205
+ # transformer compute but ranks reduce to the same final tensors). Save/push/print
1206
+ # from rank 0 only to avoid clobbering bucket files 4x and printing 4x.
1207
+ if os.environ.get("RANK", "0") == "0":
1208
+ savable = result if is_modular else _unwrap_pipeline_output(result)
1209
+ saved = _save_output(savable, self.args, self.task)
1210
+ pushed = _push_outputs(self.args, saved, self.task)
1211
+
1212
+ out.result(
1213
+ self.task,
1214
+ model=self.args.model,
1215
+ device=device,
1216
+ pipeline_class=type(pipeline).__name__,
1217
+ modular=is_modular,
1218
+ outputs=saved,
1219
+ pushed=pushed,
1220
+ seed=self.args.seed,
1221
+ output_key=self.args.output_key,
1222
+ )
1223
+ finally:
1224
+ import torch
1225
+
1226
+ if torch.distributed.is_available() and torch.distributed.is_initialized():
1227
+ torch.distributed.destroy_process_group()
diffusers/commands/schema.py ADDED
@@ -0,0 +1,287 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 The HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ """`diffusers-cli schema` — print the input schema for any pipeline repo.
16
+
17
+ Tries `DiffusionPipeline.config_name` first (so standard repos get their `__call__` signature introspected); falls back
18
+ to `ModularPipelineBlocks.from_pretrained` for modular repos. No weights are downloaded — only the small index file
19
+ (and any custom block code if `--trust-remote-code` is set).
20
+ """
21
+
22
+ from __future__ import annotations
23
+
24
+ import inspect
25
+ import re
26
+ from argparse import ArgumentParser, Namespace, _SubParsersAction
27
+ from typing import Any
28
+
29
+ from huggingface_hub.cli._output import OutputFormat, out
30
+
31
+ from ..utils import logging
32
+ from . import BaseDiffusersCLICommand
33
+
34
+
35
+ logger = logging.get_logger("diffusers-cli/schema")
36
+
37
+
38
+ def _schema(args: Namespace) -> None:
39
+ """Print the pipeline's input schema.
40
+
41
+ Tries `DiffusionPipeline.config_name` (= `model_index.json`) first; if present, introspects the declared pipeline
42
+ class's `__call__` signature. Otherwise falls back to `ModularPipelineBlocks.from_pretrained` and reads the
43
+ block-declared `inputs`. No weights downloaded either way.
44
+ """
45
+ import diffusers
46
+
47
+ try:
48
+ index = diffusers.DiffusionPipeline.load_config(args.model, token=args.token, revision=args.revision)
49
+ except OSError:
50
+ index = None
51
+
52
+ if index is not None:
53
+ class_name = index.get("_class_name")
54
+ if class_name is None:
55
+ raise SystemExit(
56
+ f"{diffusers.DiffusionPipeline.config_name} for {args.model!r} has no `_class_name` field."
57
+ )
58
+ pipeline_cls = getattr(diffusers, class_name, None)
59
+ if pipeline_cls is None:
60
+ raise SystemExit(
61
+ f"Pipeline class {class_name!r} declared in {diffusers.DiffusionPipeline.config_name} "
62
+ "is not exported by the installed diffusers."
63
+ )
64
+
65
+ sig = inspect.signature(pipeline_cls.__call__)
66
+ descriptions = _parse_docstring_args(pipeline_cls.__call__.__doc__) if args.verbose else {}
67
+ schema: list[dict[str, Any]] = []
68
+ for name, param in sig.parameters.items():
69
+ if name == "self":
70
+ continue
71
+ if param.kind in (inspect.Parameter.VAR_POSITIONAL, inspect.Parameter.VAR_KEYWORD):
72
+ continue
73
+ has_default = param.default is not inspect.Parameter.empty
74
+ schema.append(
75
+ {
76
+ "name": name,
77
+ "type_hint": str(param.annotation) if param.annotation is not inspect.Parameter.empty else None,
78
+ "default": param.default if has_default else None,
79
+ "required": not has_default,
80
+ "description": descriptions.get(name, ""),
81
+ }
82
+ )
83
+ else:
84
+ kwargs: dict[str, Any] = {"trust_remote_code": args.trust_remote_code}
85
+ if args.revision:
86
+ kwargs["revision"] = args.revision
87
+ if args.token:
88
+ kwargs["token"] = args.token
89
+
90
+ # If the repo declares custom code + external dependencies, surface them upfront so
91
+ # the user knows what to install before we hit an ImportError inside from_pretrained.
92
+ _warn_custom_block_requirements(args)
93
+
94
+ try:
95
+ blocks = diffusers.ModularPipelineBlocks.from_pretrained(args.model, **kwargs)
96
+ except Exception as e:
97
+ hint = "\nPass --trust-remote-code if it ships custom block code." if not args.trust_remote_code else ""
98
+ raise SystemExit(
99
+ f"Could not read schema for {args.model!r}: no {diffusers.DiffusionPipeline.config_name} and "
100
+ f"loading as a modular pipeline failed with:\n {type(e).__name__}: {e}{hint}"
101
+ ) from e
102
+
103
+ class_name = type(blocks).__name__
104
+ schema = [
105
+ {
106
+ "name": p.name,
107
+ "type_hint": str(p.type_hint) if p.type_hint is not None else None,
108
+ "default": p.default,
109
+ "required": p.required,
110
+ "description": p.description,
111
+ }
112
+ for p in blocks.inputs
113
+ ]
114
+
115
+ if out.mode == OutputFormat.json:
116
+ out.dict({"task": "schema", "model": args.model, "pipeline_class": class_name, "inputs": schema})
117
+ elif out.mode == OutputFormat.agent:
118
+ out.table(schema, headers=["name", "required", "type_hint", "default", "description"])
119
+ else:
120
+ out.text(f"{class_name} ({args.model}) inputs:")
121
+ for entry in schema:
122
+ tag = "required" if entry["required"] else f"optional, default={entry['default']!r}"
123
+ out.text(f" {entry['name']} ({tag})")
124
+ if entry["type_hint"]:
125
+ out.text(f" type: {entry['type_hint']}")
126
+ if entry["description"]:
127
+ out.text(f" desc: {entry['description']}")
128
+
129
+
130
+ def _warn_custom_block_requirements(args: Namespace) -> None:
131
+ """Warn upfront when a modular block ships custom code with declared external dependencies.
132
+
133
+ Reads `modular_config.json` if present; if it has an `auto_map` (custom code) and a non-empty `requirements`
134
+ list/dict, prints a heads-up. `from_pretrained` will otherwise fail with an `ImportError` deep in the loader stack
135
+ when a listed dep is missing.
136
+ """
137
+ import diffusers
138
+
139
+ try:
140
+ config = diffusers.ModularPipelineBlocks.load_config(args.model, token=args.token, revision=args.revision)
141
+ except Exception:
142
+ return # no modular_config.json or unreachable — nothing to warn about
143
+ if not isinstance(config, dict):
144
+ return
145
+ if not config.get("auto_map"):
146
+ return
147
+ requirements = config.get("requirements")
148
+ if not requirements:
149
+ return
150
+
151
+ # `requirements` may be a dict {name: version} or (older repos) a list of [name, version] pairs.
152
+ if isinstance(requirements, dict):
153
+ pairs = list(requirements.items())
154
+ elif isinstance(requirements, list):
155
+ pairs = [(item[0], item[1]) for item in requirements if isinstance(item, (list, tuple)) and len(item) >= 2]
156
+ else:
157
+ pairs = []
158
+ if not pairs:
159
+ return
160
+
161
+ formatted = ", ".join(f"{name}=={version}" for name, version in pairs)
162
+ logger.warning(
163
+ f"{args.model!r} ships custom block code with external dependencies: {formatted}. "
164
+ "You will need to install these in order to determine the pipeline schema."
165
+ )
166
+
167
+
168
+ def _parse_docstring_args(docstring: str | None) -> dict[str, str]:
169
+ """Extract per-argument descriptions from a Google-style `Args:` block.
170
+
171
+ Returns a `{name: description}` mapping. Best-effort — unrecognised formats just yield an empty dict rather than
172
+ raising.
173
+ """
174
+ if not docstring:
175
+ return {}
176
+
177
+ lines = docstring.expandtabs().splitlines()
178
+ start = None
179
+ section_indent = 0
180
+ for i, line in enumerate(lines):
181
+ if line.strip() in ("Args:", "Arguments:", "Parameters:"):
182
+ start = i + 1
183
+ section_indent = len(line) - len(line.lstrip())
184
+ break
185
+ if start is None:
186
+ return {}
187
+
188
+ descriptions: dict[str, str] = {}
189
+ current_name: str | None = None
190
+ current_lines: list[str] = []
191
+ arg_indent: int | None = None
192
+ name_pattern = re.compile(r"^(\w+)\s*(?:\([^)]*\))?\s*:?\s*(.*)$")
193
+
194
+ def _flush() -> None:
195
+ if current_name and current_lines:
196
+ descriptions[current_name] = " ".join(s.strip() for s in current_lines).strip()
197
+
198
+ for line in lines[start:]:
199
+ if not line.strip():
200
+ continue
201
+ indent = len(line) - len(line.lstrip())
202
+ # A new top-level section ends the Args block.
203
+ if indent <= section_indent and line.strip().endswith(":"):
204
+ break
205
+ if arg_indent is None:
206
+ arg_indent = indent
207
+ if indent == arg_indent:
208
+ _flush()
209
+ current_lines = []
210
+ match = name_pattern.match(line.strip())
211
+ if match:
212
+ current_name = match.group(1)
213
+ tail = match.group(2).strip()
214
+ if tail:
215
+ current_lines.append(tail)
216
+ else:
217
+ current_name = None
218
+ elif current_name is not None and indent > arg_indent:
219
+ current_lines.append(line.strip())
220
+ _flush()
221
+ return descriptions
222
+
223
+
224
+ class SchemaCommand(BaseDiffusersCLICommand):
225
+ task = "schema"
226
+
227
+ @staticmethod
228
+ def register_subcommand(subparsers: _SubParsersAction) -> None:
229
+ from argparse import RawDescriptionHelpFormatter
230
+
231
+ epilog = (
232
+ "Examples\n"
233
+ " $ diffusers-cli schema -m stabilityai/stable-diffusion-xl-base-1.0\n"
234
+ " $ diffusers-cli schema -m black-forest-labs/FLUX.1-dev --verbose\n"
235
+ " $ diffusers-cli --format json schema -m stabilityai/stable-diffusion-xl-base-1.0\n"
236
+ "\n"
237
+ "Learn more\n"
238
+ " Use `diffusers-cli <command> --help` for more information about a command.\n"
239
+ " Read the documentation at https://huggingface.co/docs/diffusers\n"
240
+ )
241
+
242
+ parser: ArgumentParser = subparsers.add_parser(
243
+ "schema",
244
+ help="Print the input schema for a diffusers pipeline repo. No weights downloaded.",
245
+ usage="\n diffusers-cli schema [options]",
246
+ epilog=epilog,
247
+ formatter_class=RawDescriptionHelpFormatter,
248
+ )
249
+ parser._optionals.title = "Options"
250
+ parser.add_argument(
251
+ "--model",
252
+ "-m",
253
+ required=True,
254
+ help="Model id on the Hugging Face Hub or local path.",
255
+ )
256
+ parser.add_argument(
257
+ "--revision",
258
+ default=None,
259
+ help="Model revision (branch, tag, or commit SHA).",
260
+ )
261
+ parser.add_argument(
262
+ "--token",
263
+ default=None,
264
+ help="Hugging Face token for gated/private models.",
265
+ )
266
+ parser.add_argument(
267
+ "--trust-remote-code",
268
+ action="store_true",
269
+ help="Allow custom code from the Hub (required for modular pipelines that ship block code).",
270
+ )
271
+ parser.add_argument(
272
+ "--verbose",
273
+ "-v",
274
+ action="store_true",
275
+ help=(
276
+ "Also include per-argument descriptions from the pipeline's __call__ docstring. "
277
+ "Modular pipelines always include block-declared descriptions; --verbose populates "
278
+ "the equivalent field for standard pipelines by parsing the Google-style Args: block."
279
+ ),
280
+ )
281
+ parser.set_defaults(func=SchemaCommand)
282
+
283
+ def __init__(self, args: Namespace):
284
+ self.args = args
285
+
286
+ def run(self) -> None:
287
+ _schema(self.args)
diffusers/commands/skills.py ADDED
@@ -0,0 +1,344 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 The HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+ """`diffusers-cli skills` — install Agent Skills bundles.
15
+
16
+ Skill bundles live under `.ai/skills/<name>/` in the diffusers repo and follow the Agent Skills standard: a directory
17
+ containing `SKILL.md` (plus optional resources). Installs to `.agents/skills/<name>/` which Claude, Codex, and Cursor
18
+ all discover.
19
+ """
20
+
21
+ from __future__ import annotations
22
+
23
+ import os
24
+ import shutil
25
+ from argparse import ArgumentParser, Namespace, _SubParsersAction
26
+ from pathlib import Path
27
+
28
+ import httpx
29
+ from huggingface_hub.cli._output import out
30
+
31
+ from ..utils import logging
32
+ from ..utils.constants import DIFFUSERS_REQUEST_TIMEOUT
33
+ from . import BaseDiffusersCLICommand
34
+
35
+
36
+ logger = logging.get_logger("diffusers-cli/skills")
37
+
38
+
39
+ _REGISTRY_BASE = "https://api.github.com/repos/huggingface/diffusers/contents/.ai/skills"
40
+ _REGISTRY_REF = "main"
41
+
42
+ # Native skill-discovery paths per agent. Claude Code reads only `.claude/skills/`; Codex and
43
+ # Cursor read `.agents/skills/` (Cursor also honors `.claude/skills/` via compat, but installing
44
+ # to `.agents/skills/` is the portable choice for both).
45
+ _CLAUDE_SKILLS_DIR = Path(".claude") / "skills"
46
+ _AGENTS_SKILLS_DIR = Path(".agents") / "skills"
47
+
48
+ # Env vars set by each agent when it launches the CLI. Values are the install path to use.
49
+ _AGENT_ENV_TO_DIR: dict[str, Path] = {
50
+ "CLAUDECODE": _CLAUDE_SKILLS_DIR,
51
+ "CLAUDE_CODE": _CLAUDE_SKILLS_DIR,
52
+ "CODEX_SANDBOX": _AGENTS_SKILLS_DIR,
53
+ "CURSOR_AI": _AGENTS_SKILLS_DIR,
54
+ }
55
+ # When no agent env var is set, install to every native path so whichever agent the user
56
+ # later switches to picks the skill up.
57
+ _ALL_INSTALL_DIRS: tuple[Path, ...] = (_CLAUDE_SKILLS_DIR, _AGENTS_SKILLS_DIR)
58
+
59
+ # Empty marker dropped inside each installed skill dir so `update` can distinguish our
60
+ # installs from user-placed skills at the same paths.
61
+ _MANAGED_MARKER_FILE = ".diffusers-skill-managed"
62
+
63
+
64
+ # ---------------------------------------------------------------------------
65
+ # Registry fetch
66
+ # ---------------------------------------------------------------------------
67
+
68
+
69
+ def _registry_url(name: str = "") -> str:
70
+ """API URL for the registry root, or for a single skill bundle when `name` is given."""
71
+ path = f"/{name}" if name else ""
72
+ return f"{_REGISTRY_BASE}{path}?ref={_REGISTRY_REF}"
73
+
74
+
75
+ def _fetch_json(url: str) -> list[dict]:
76
+ try:
77
+ resp = httpx.get(url, timeout=DIFFUSERS_REQUEST_TIMEOUT)
78
+ resp.raise_for_status()
79
+ return resp.json()
80
+ except httpx.HTTPStatusError as e:
81
+ if e.response.status_code == 404:
82
+ raise SystemExit(f"Not found in registry: {url}") from e
83
+ raise SystemExit(f"Registry fetch failed: HTTP {e.response.status_code} {e.response.reason_phrase}") from e
84
+ except httpx.HTTPError as e:
85
+ raise SystemExit(f"Could not reach registry: {e}") from e
86
+
87
+
88
+ def _walk_skill_files(name: str) -> list[tuple[str, str]]:
89
+ files: list[tuple[str, str]] = []
90
+
91
+ def _walk(api_url: str, prefix: str) -> None:
92
+ for entry in _fetch_json(api_url):
93
+ if entry["type"] == "file":
94
+ files.append((f"{prefix}{entry['name']}", entry["download_url"]))
95
+ elif entry["type"] == "dir":
96
+ _walk(entry["url"], f"{prefix}{entry['name']}/")
97
+
98
+ _walk(_registry_url(name), "")
99
+ return files
100
+
101
+
102
+ def _download_skill_bundle(name: str) -> dict[str, bytes]:
103
+ files = _walk_skill_files(name)
104
+ if not files:
105
+ raise SystemExit(f"Skill '{name}' has no files in the registry.")
106
+ bundle: dict[str, bytes] = {}
107
+ for rel_path, url in files:
108
+ resp = httpx.get(url, timeout=DIFFUSERS_REQUEST_TIMEOUT)
109
+ resp.raise_for_status()
110
+ bundle[rel_path] = resp.content
111
+ return bundle
112
+
113
+
114
+ # ---------------------------------------------------------------------------
115
+ # Install / discovery
116
+ # ---------------------------------------------------------------------------
117
+
118
+
119
+ def _detect_install_dirs() -> tuple[Path, ...]:
120
+ """Pick where to install based on the launching agent.
121
+
122
+ If we detect a specific agent from its env var, install only there. If nothing is detected, install to every native
123
+ path so any agent picks the skill up later.
124
+ """
125
+ for env_var, skills_dir in _AGENT_ENV_TO_DIR.items():
126
+ if os.environ.get(env_var):
127
+ return (skills_dir,)
128
+ return _ALL_INSTALL_DIRS
129
+
130
+
131
+ def _install_skill(name: str, bundle: dict[str, bytes], root: Path, skills_dir: Path, force: bool) -> Path:
132
+ skill_dir = root / skills_dir / name
133
+ if skill_dir.exists():
134
+ if not force:
135
+ raise SystemExit(f"Skill already installed at {skill_dir}. Use --force to reinstall.")
136
+ shutil.rmtree(skill_dir)
137
+ skill_dir.mkdir(parents=True, exist_ok=True)
138
+ for rel_path, data in bundle.items():
139
+ target = skill_dir / rel_path
140
+ target.parent.mkdir(parents=True, exist_ok=True)
141
+ target.write_bytes(data)
142
+ (skill_dir / _MANAGED_MARKER_FILE).touch()
143
+ return skill_dir
144
+
145
+
146
+ def _has_local_changes(skill_dir: Path, bundle: dict[str, bytes]) -> bool:
147
+ """True if the installed skill has any file that differs from `bundle` or has extra files.
148
+
149
+ The marker file is ignored. Compares raw bytes so a whitespace-only edit still counts as dirty.
150
+ """
151
+ on_disk: dict[str, bytes] = {}
152
+ for path in skill_dir.rglob("*"):
153
+ if not path.is_file():
154
+ continue
155
+ rel = str(path.relative_to(skill_dir))
156
+ if rel == _MANAGED_MARKER_FILE:
157
+ continue
158
+ on_disk[rel] = path.read_bytes()
159
+ return on_disk != bundle
160
+
161
+
162
+ def _discover_installed(root: Path) -> list[tuple[Path, str]]:
163
+ """Return `(skills_dir, name)` pairs for every managed install under `root`."""
164
+ found: list[tuple[Path, str]] = []
165
+ for skills_dir in _ALL_INSTALL_DIRS:
166
+ skills_root = root / skills_dir
167
+ if not skills_root.exists():
168
+ continue
169
+ for d in sorted(skills_root.iterdir()):
170
+ if d.is_dir() and (d / _MANAGED_MARKER_FILE).exists():
171
+ found.append((skills_dir, d.name))
172
+ return found
173
+
174
+
175
+ class SkillsCommand(BaseDiffusersCLICommand):
176
+ @staticmethod
177
+ def register_subcommand(subparsers: _SubParsersAction) -> None:
178
+ parser: ArgumentParser = subparsers.add_parser(
179
+ "skills",
180
+ help="Manage Agent Skills for AI assistants.",
181
+ usage="\n diffusers-cli skills <add|list|update|preview> [options]",
182
+ )
183
+ parser._optionals.title = "Options"
184
+ actions = parser.add_subparsers(dest="skills_action", required=True, metavar="<action>")
185
+
186
+ add = actions.add_parser("add", help="Download and install a skill.")
187
+ add.add_argument(
188
+ "name",
189
+ nargs="?",
190
+ default=None,
191
+ help="Skill name (e.g. diffusers-cli, custom-blocks). Omit and pass --all to install every skill.",
192
+ )
193
+ add.add_argument(
194
+ "--all",
195
+ dest="install_all",
196
+ action="store_true",
197
+ help="Install every skill in the registry. Mutually exclusive with a positional name.",
198
+ )
199
+ add.add_argument(
200
+ "--global",
201
+ "-g",
202
+ dest="install_global",
203
+ action="store_true",
204
+ help="Install globally (user-level) instead of in the current project directory.",
205
+ )
206
+ add.add_argument("--force", action="store_true", help="Overwrite existing skills in the destination.")
207
+ add.set_defaults(func=SkillsCommand)
208
+
209
+ list_action = actions.add_parser("list", help="List available skills in the registry.")
210
+ list_action.set_defaults(func=SkillsCommand)
211
+
212
+ update = actions.add_parser("update", help="Re-download and reinstall managed skills.")
213
+ update.add_argument(
214
+ "name",
215
+ nargs="?",
216
+ default=None,
217
+ help="Optional installed skill name to update. Omit to update every managed skill.",
218
+ )
219
+ update.add_argument(
220
+ "--global",
221
+ "-g",
222
+ dest="install_global",
223
+ action="store_true",
224
+ help="Update skills installed globally (user-level) instead of the current project.",
225
+ )
226
+ update.add_argument(
227
+ "--force",
228
+ action="store_true",
229
+ help="Overwrite skills even if they have local modifications since install.",
230
+ )
231
+ update.set_defaults(func=SkillsCommand)
232
+
233
+ preview = actions.add_parser("preview", help="Print a skill's SKILL.md from the registry.")
234
+ preview.add_argument("name", help="Skill name to preview.")
235
+ preview.set_defaults(func=SkillsCommand)
236
+
237
+ def __init__(self, args: Namespace):
238
+ self.args = args
239
+
240
+ def run(self) -> None:
241
+ if self.args.skills_action == "add":
242
+ self._add()
243
+ elif self.args.skills_action == "list":
244
+ self._list()
245
+ elif self.args.skills_action == "update":
246
+ self._update()
247
+ elif self.args.skills_action == "preview":
248
+ self._preview()
249
+
250
+ def _add(self) -> None:
251
+ if self.args.install_all and self.args.name:
252
+ raise SystemExit("--all and a positional skill name are mutually exclusive.")
253
+ if not self.args.install_all and not self.args.name:
254
+ raise SystemExit("Pass a skill name (e.g. diffusers-cli) or --all to install every skill.")
255
+
256
+ root = Path.home() if self.args.install_global else Path.cwd()
257
+ install_dirs = _detect_install_dirs()
258
+ names = self._resolve_names()
259
+
260
+ installed: list[str] = []
261
+ failed: list[str] = []
262
+ for name in names:
263
+ try:
264
+ bundle = _download_skill_bundle(name)
265
+ for skills_dir in install_dirs:
266
+ _install_skill(name, bundle, root, skills_dir, self.args.force)
267
+ installed.append(name)
268
+ except (SystemExit, httpx.HTTPError) as e:
269
+ # Downgrade to a warning so one broken skill doesn't abort the batch.
270
+ logger.warning(f"Skipping skill {name!r}: {e}")
271
+ failed.append(name)
272
+
273
+ if not installed:
274
+ raise SystemExit(f"No skills installed. Failed: {failed}")
275
+ out.result(
276
+ f"Installed {len(installed)} skill(s)",
277
+ installed=", ".join(installed),
278
+ failed=", ".join(failed) if failed else None,
279
+ paths=", ".join(str(root / d) for d in install_dirs),
280
+ )
281
+
282
+ def _update(self) -> None:
283
+ root = Path.home() if self.args.install_global else Path.cwd()
284
+ installed = _discover_installed(root)
285
+ if self.args.name is not None:
286
+ installed = [entry for entry in installed if entry[1] == self.args.name]
287
+ if not installed:
288
+ raise SystemExit(f"No installed skill named {self.args.name!r} found under {root}.")
289
+ if not installed:
290
+ raise SystemExit(f"No managed skills found under {root}.")
291
+
292
+ # Group by skill name so we redownload each bundle once even if it's installed to
293
+ # multiple locations (e.g. both .claude/skills/ and .agents/skills/).
294
+ by_name: dict[str, list[Path]] = {}
295
+ for skills_dir, name in installed:
296
+ by_name.setdefault(name, []).append(skills_dir)
297
+
298
+ updated: list[str] = []
299
+ failed: list[str] = []
300
+ skipped: list[str] = []
301
+ for name, dirs in sorted(by_name.items()):
302
+ try:
303
+ bundle = _download_skill_bundle(name)
304
+ for skills_dir in dirs:
305
+ skill_dir = root / skills_dir / name
306
+ if not self.args.force and _has_local_changes(skill_dir, bundle):
307
+ logger.warning(
308
+ f"Skill {name!r} at {skill_dir} has local modifications; "
309
+ "skipping. Pass --force to overwrite them."
310
+ )
311
+ skipped.append(name)
312
+ continue
313
+ _install_skill(name, bundle, root, skills_dir, force=True)
314
+ updated.append(name)
315
+ except (SystemExit, httpx.HTTPError) as e:
316
+ logger.warning(f"Skipping skill {name!r}: {e}")
317
+ failed.append(name)
318
+
319
+ out.result(
320
+ f"Updated {len(updated)} skill(s)",
321
+ updated=", ".join(updated),
322
+ skipped=", ".join(skipped) if skipped else None,
323
+ failed=", ".join(failed) if failed else None,
324
+ )
325
+
326
+ def _preview(self) -> None:
327
+ bundle = _download_skill_bundle(self.args.name)
328
+ skill_md = bundle.get("SKILL.md")
329
+ if skill_md is None:
330
+ raise SystemExit(f"Skill {self.args.name!r} has no SKILL.md in the registry.")
331
+ print(skill_md.decode())
332
+
333
+ def _list(self) -> None:
334
+ entries = _fetch_json(_registry_url())
335
+ skills = [{"name": e["name"]} for e in entries if e["type"] == "dir" and not e["name"].startswith(".")]
336
+ if not skills:
337
+ raise SystemExit("No skills found in registry.")
338
+ out.table(skills, headers=["name"])
339
+
340
+ def _resolve_names(self) -> list[str]:
341
+ if self.args.install_all:
342
+ entries = _fetch_json(_registry_url())
343
+ return sorted(e["name"] for e in entries if e["type"] == "dir" and not e["name"].startswith("."))
344
+ return [self.args.name]
diffusers/configuration_utils.py ADDED
@@ -0,0 +1,752 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # coding=utf-8
2
+ # Copyright 2025 The HuggingFace Inc. team.
3
+ # Copyright (c) 2022, NVIDIA CORPORATION. All rights reserved.
4
+ #
5
+ # Licensed under the Apache License, Version 2.0 (the "License");
6
+ # you may not use this file except in compliance with the License.
7
+ # You may obtain a copy of the License at
8
+ #
9
+ # http://www.apache.org/licenses/LICENSE-2.0
10
+ #
11
+ # Unless required by applicable law or agreed to in writing, software
12
+ # distributed under the License is distributed on an "AS IS" BASIS,
13
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
14
+ # See the License for the specific language governing permissions and
15
+ # limitations under the License.
16
+ """ConfigMixin base class and utilities."""
17
+
18
+ import functools
19
+ import importlib
20
+ import inspect
21
+ import json
22
+ import os
23
+ import re
24
+ from collections import OrderedDict
25
+ from pathlib import Path
26
+ from typing import Any
27
+
28
+ import numpy as np
29
+ from huggingface_hub import DDUFEntry, create_repo, hf_hub_download
30
+ from huggingface_hub.utils import (
31
+ EntryNotFoundError,
32
+ HfHubHTTPError,
33
+ RepositoryNotFoundError,
34
+ RevisionNotFoundError,
35
+ validate_hf_hub_args,
36
+ )
37
+ from typing_extensions import Self
38
+
39
+ from . import __version__
40
+ from .utils import (
41
+ HUGGINGFACE_CO_RESOLVE_ENDPOINT,
42
+ DummyObject,
43
+ deprecate,
44
+ extract_commit_hash,
45
+ http_user_agent,
46
+ logging,
47
+ )
48
+
49
+
50
+ logger = logging.get_logger(__name__)
51
+
52
+ _re_configuration_file = re.compile(r"config\.(.*)\.json")
53
+
54
+
55
+ class FrozenDict(OrderedDict):
56
+ def __init__(self, *args, **kwargs):
57
+ super().__init__(*args, **kwargs)
58
+
59
+ for key, value in self.items():
60
+ setattr(self, key, value)
61
+
62
+ self.__frozen = True
63
+
64
+ def __delitem__(self, *args, **kwargs):
65
+ raise Exception(f"You cannot use ``__delitem__`` on a {self.__class__.__name__} instance.")
66
+
67
+ def setdefault(self, *args, **kwargs):
68
+ raise Exception(f"You cannot use ``setdefault`` on a {self.__class__.__name__} instance.")
69
+
70
+ def pop(self, *args, **kwargs):
71
+ raise Exception(f"You cannot use ``pop`` on a {self.__class__.__name__} instance.")
72
+
73
+ def update(self, *args, **kwargs):
74
+ raise Exception(f"You cannot use ``update`` on a {self.__class__.__name__} instance.")
75
+
76
+ def __setattr__(self, name, value):
77
+ if hasattr(self, "__frozen") and self.__frozen:
78
+ raise Exception(f"You cannot use ``__setattr__`` on a {self.__class__.__name__} instance.")
79
+ super().__setattr__(name, value)
80
+
81
+ def __setitem__(self, name, value):
82
+ if hasattr(self, "__frozen") and self.__frozen:
83
+ raise Exception(f"You cannot use ``__setattr__`` on a {self.__class__.__name__} instance.")
84
+ super().__setitem__(name, value)
85
+
86
+
87
+ class ConfigMixin:
88
+ r"""
89
+ Base class for all configuration classes. All configuration parameters are stored under `self.config`. Also
90
+ provides the [`~ConfigMixin.from_config`] and [`~ConfigMixin.save_config`] methods for loading, downloading, and
91
+ saving classes that inherit from [`ConfigMixin`].
92
+
93
+ Class attributes:
94
+ - **config_name** (`str`) -- A filename under which the config should stored when calling
95
+ [`~ConfigMixin.save_config`] (should be overridden by parent class).
96
+ - **ignore_for_config** (`list[str]`) -- A list of attributes that should not be saved in the config (should be
97
+ overridden by subclass).
98
+ - **has_compatibles** (`bool`) -- Whether the class has compatible classes (should be overridden by subclass).
99
+ - **_deprecated_kwargs** (`list[str]`) -- Keyword arguments that are deprecated. Note that the `init` function
100
+ should only have a `kwargs` argument if at least one argument is deprecated (should be overridden by
101
+ subclass).
102
+ """
103
+
104
+ config_name = None
105
+ ignore_for_config = []
106
+ has_compatibles = False
107
+
108
+ _deprecated_kwargs = []
109
+ _auto_class = None
110
+
111
+ @classmethod
112
+ def register_for_auto_class(cls, auto_class="AutoModel"):
113
+ """
114
+ Register this class with the given auto class so that it can be loaded with `AutoModel.from_pretrained(...,
115
+ trust_remote_code=True)`.
116
+
117
+ When the config is saved, the resulting `config.json` will include an `auto_map` entry mapping the auto class
118
+ to this class's module and class name.
119
+
120
+ Args:
121
+ auto_class (`str` or type, *optional*, defaults to `"AutoModel"`):
122
+ The auto class to register this class with. Can be a string (e.g. `"AutoModel"`) or the class itself.
123
+ Currently only `"AutoModel"` is supported.
124
+
125
+ Example:
126
+
127
+ ```python
128
+ from diffusers import ModelMixin, ConfigMixin
129
+
130
+
131
+ class MyCustomModel(ModelMixin, ConfigMixin): ...
132
+
133
+
134
+ MyCustomModel.register_for_auto_class("AutoModel")
135
+ ```
136
+ """
137
+ if auto_class != "AutoModel":
138
+ raise ValueError(f"Only 'AutoModel' is supported, got '{auto_class}'.")
139
+
140
+ cls._auto_class = auto_class
141
+
142
+ def register_to_config(self, **kwargs):
143
+ if self.config_name is None:
144
+ raise NotImplementedError(f"Make sure that {self.__class__} has defined a class name `config_name`")
145
+ # Special case for `kwargs` used in deprecation warning added to schedulers
146
+ # TODO: remove this when we remove the deprecation warning, and the `kwargs` argument,
147
+ # or solve in a more general way.
148
+ kwargs.pop("kwargs", None)
149
+
150
+ if not hasattr(self, "_internal_dict"):
151
+ internal_dict = kwargs
152
+ else:
153
+ previous_dict = dict(self._internal_dict)
154
+ internal_dict = {**self._internal_dict, **kwargs}
155
+ logger.debug(f"Updating config from {previous_dict} to {internal_dict}")
156
+
157
+ self._internal_dict = FrozenDict(internal_dict)
158
+
159
+ def __getattr__(self, name: str) -> Any:
160
+ """The only reason we overwrite `getattr` here is to gracefully deprecate accessing
161
+ config attributes directly. See https://github.com/huggingface/diffusers/pull/3129
162
+
163
+ This function is mostly copied from PyTorch's __getattr__ overwrite:
164
+ https://pytorch.org/docs/stable/_modules/torch/nn/modules/module.html#Module
165
+ """
166
+
167
+ is_in_config = "_internal_dict" in self.__dict__ and hasattr(self.__dict__["_internal_dict"], name)
168
+ is_attribute = name in self.__dict__
169
+
170
+ if is_in_config and not is_attribute:
171
+ deprecation_message = f"Accessing config attribute `{name}` directly via '{type(self).__name__}' object attribute is deprecated. Please access '{name}' over '{type(self).__name__}'s config object instead, e.g. 'scheduler.config.{name}'."
172
+ deprecate("direct config name access", "1.0.0", deprecation_message, standard_warn=False)
173
+ return self._internal_dict[name]
174
+
175
+ raise AttributeError(f"'{type(self).__name__}' object has no attribute '{name}'")
176
+
177
+ def save_config(self, save_directory: str | os.PathLike, push_to_hub: bool = False, **kwargs):
178
+ """
179
+ Save a configuration object to the directory specified in `save_directory` so that it can be reloaded using the
180
+ [`~ConfigMixin.from_config`] class method.
181
+
182
+ Args:
183
+ save_directory (`str` or `os.PathLike`):
184
+ Directory where the configuration JSON file is saved (will be created if it does not exist).
185
+ push_to_hub (`bool`, *optional*, defaults to `False`):
186
+ Whether or not to push your model to the Hugging Face Hub after saving it. You can specify the
187
+ repository you want to push to with `repo_id` (will default to the name of `save_directory` in your
188
+ namespace).
189
+ kwargs (`dict[str, Any]`, *optional*):
190
+ Additional keyword arguments passed along to the [`~utils.PushToHubMixin.push_to_hub`] method.
191
+ """
192
+ if os.path.isfile(save_directory):
193
+ raise AssertionError(f"Provided path ({save_directory}) should be a directory, not a file")
194
+
195
+ os.makedirs(save_directory, exist_ok=True)
196
+
197
+ # If we save using the predefined names, we can load using `from_config`
198
+ output_config_file = os.path.join(save_directory, self.config_name)
199
+
200
+ self.to_json_file(output_config_file)
201
+ logger.info(f"Configuration saved in {output_config_file}")
202
+
203
+ if push_to_hub:
204
+ commit_message = kwargs.pop("commit_message", None)
205
+ private = kwargs.pop("private", None)
206
+ create_pr = kwargs.pop("create_pr", False)
207
+ token = kwargs.pop("token", None)
208
+ repo_id = kwargs.pop("repo_id", save_directory.split(os.path.sep)[-1])
209
+ repo_id = create_repo(repo_id, exist_ok=True, private=private, token=token).repo_id
210
+ subfolder = kwargs.pop("subfolder", None)
211
+
212
+ self._upload_folder(
213
+ save_directory,
214
+ repo_id,
215
+ token=token,
216
+ commit_message=commit_message,
217
+ create_pr=create_pr,
218
+ subfolder=subfolder,
219
+ )
220
+
221
+ @classmethod
222
+ def from_config(
223
+ cls, config: FrozenDict | dict[str, Any] = None, return_unused_kwargs=False, **kwargs
224
+ ) -> Self | tuple[Self, dict[str, Any]]:
225
+ r"""
226
+ Instantiate a Python class from a config dictionary.
227
+
228
+ Parameters:
229
+ config (`dict[str, Any]`):
230
+ A config dictionary from which the Python class is instantiated. Make sure to only load configuration
231
+ files of compatible classes.
232
+ return_unused_kwargs (`bool`, *optional*, defaults to `False`):
233
+ Whether kwargs that are not consumed by the Python class should be returned or not.
234
+ kwargs (remaining dictionary of keyword arguments, *optional*):
235
+ Can be used to update the configuration object (after it is loaded) and initiate the Python class.
236
+ `**kwargs` are passed directly to the underlying scheduler/model's `__init__` method and eventually
237
+ overwrite the same named arguments in `config`.
238
+
239
+ Returns:
240
+ [`ModelMixin`] or [`SchedulerMixin`]:
241
+ A model or scheduler object instantiated from a config dictionary.
242
+
243
+ Examples:
244
+
245
+ ```python
246
+ >>> from diffusers import DDPMScheduler, DDIMScheduler, PNDMScheduler
247
+
248
+ >>> # Download scheduler from huggingface.co and cache.
249
+ >>> scheduler = DDPMScheduler.from_pretrained("google/ddpm-cifar10-32")
250
+
251
+ >>> # Instantiate DDIM scheduler class with same config as DDPM
252
+ >>> scheduler = DDIMScheduler.from_config(scheduler.config)
253
+
254
+ >>> # Instantiate PNDM scheduler class with same config as DDPM
255
+ >>> scheduler = PNDMScheduler.from_config(scheduler.config)
256
+ ```
257
+ """
258
+ # <===== TO BE REMOVED WITH DEPRECATION
259
+ # TODO(Patrick) - make sure to remove the following lines when config=="model_path" is deprecated
260
+ if "pretrained_model_name_or_path" in kwargs:
261
+ config = kwargs.pop("pretrained_model_name_or_path")
262
+
263
+ if config is None:
264
+ raise ValueError("Please make sure to provide a config as the first positional argument.")
265
+ # ======>
266
+
267
+ if not isinstance(config, dict):
268
+ deprecation_message = "It is deprecated to pass a pretrained model name or path to `from_config`."
269
+ if "Scheduler" in cls.__name__:
270
+ deprecation_message += (
271
+ f"If you were trying to load a scheduler, please use {cls}.from_pretrained(...) instead."
272
+ " Otherwise, please make sure to pass a configuration dictionary instead. This functionality will"
273
+ " be removed in v1.0.0."
274
+ )
275
+ elif "Model" in cls.__name__:
276
+ deprecation_message += (
277
+ f"If you were trying to load a model, please use {cls}.load_config(...) followed by"
278
+ f" {cls}.from_config(...) instead. Otherwise, please make sure to pass a configuration dictionary"
279
+ " instead. This functionality will be removed in v1.0.0."
280
+ )
281
+ deprecate("config-passed-as-path", "1.0.0", deprecation_message, standard_warn=False)
282
+ config, kwargs = cls.load_config(pretrained_model_name_or_path=config, return_unused_kwargs=True, **kwargs)
283
+
284
+ init_dict, unused_kwargs, hidden_dict = cls.extract_init_dict(config, **kwargs)
285
+
286
+ # Allow dtype to be specified on initialization
287
+ if "dtype" in unused_kwargs:
288
+ init_dict["dtype"] = unused_kwargs.pop("dtype")
289
+
290
+ # add possible deprecated kwargs
291
+ for deprecated_kwarg in cls._deprecated_kwargs:
292
+ if deprecated_kwarg in unused_kwargs:
293
+ init_dict[deprecated_kwarg] = unused_kwargs.pop(deprecated_kwarg)
294
+
295
+ # Return model and optionally state and/or unused_kwargs
296
+ model = cls(**init_dict)
297
+
298
+ # make sure to also save config parameters that might be used for compatible classes
299
+ # update _class_name
300
+ if "_class_name" in hidden_dict:
301
+ hidden_dict["_class_name"] = cls.__name__
302
+
303
+ model.register_to_config(**hidden_dict)
304
+
305
+ # add hidden kwargs of compatible classes to unused_kwargs
306
+ unused_kwargs = {**unused_kwargs, **hidden_dict}
307
+
308
+ if return_unused_kwargs:
309
+ return (model, unused_kwargs)
310
+ else:
311
+ return model
312
+
313
+ @classmethod
314
+ def get_config_dict(cls, *args, **kwargs):
315
+ deprecation_message = (
316
+ f" The function get_config_dict is deprecated. Please use {cls}.load_config instead. This function will be"
317
+ " removed in version v1.0.0"
318
+ )
319
+ deprecate("get_config_dict", "1.0.0", deprecation_message, standard_warn=False)
320
+ return cls.load_config(*args, **kwargs)
321
+
322
+ @classmethod
323
+ @validate_hf_hub_args
324
+ def load_config(
325
+ cls,
326
+ pretrained_model_name_or_path: str | os.PathLike,
327
+ return_unused_kwargs=False,
328
+ return_commit_hash=False,
329
+ **kwargs,
330
+ ) -> tuple[dict[str, Any], dict[str, Any]]:
331
+ r"""
332
+ Load a model or scheduler configuration.
333
+
334
+ Parameters:
335
+ pretrained_model_name_or_path (`str` or `os.PathLike`, *optional*):
336
+ Can be either:
337
+
338
+ - A string, the *model id* (for example `google/ddpm-celebahq-256`) of a pretrained model hosted on
339
+ the Hub.
340
+ - A path to a *directory* (for example `./my_model_directory`) containing model weights saved with
341
+ [`~ConfigMixin.save_config`].
342
+
343
+ cache_dir (`str | os.PathLike`, *optional*):
344
+ Path to a directory where a downloaded pretrained model configuration is cached if the standard cache
345
+ is not used.
346
+ force_download (`bool`, *optional*, defaults to `False`):
347
+ Whether or not to force the (re-)download of the model weights and configuration files, overriding the
348
+ cached versions if they exist.
349
+ proxies (`dict[str, str]`, *optional*):
350
+ A dictionary of proxy servers to use by protocol or endpoint, for example, `{'http': 'foo.bar:3128',
351
+ 'http://hostname': 'foo.bar:4012'}`. The proxies are used on each request.
352
+ output_loading_info(`bool`, *optional*, defaults to `False`):
353
+ Whether or not to also return a dictionary containing missing keys, unexpected keys and error messages.
354
+ local_files_only (`bool`, *optional*, defaults to `False`):
355
+ Whether to only load local model weights and configuration files or not. If set to `True`, the model
356
+ won't be downloaded from the Hub.
357
+ token (`str` or *bool*, *optional*):
358
+ The token to use as HTTP bearer authorization for remote files. If `True`, the token generated from
359
+ `diffusers-cli login` (stored in `~/.huggingface`) is used.
360
+ revision (`str`, *optional*, defaults to `"main"`):
361
+ The specific model version to use. It can be a branch name, a tag name, a commit id, or any identifier
362
+ allowed by Git.
363
+ subfolder (`str`, *optional*, defaults to `""`):
364
+ The subfolder location of a model file within a larger model repository on the Hub or locally.
365
+ return_unused_kwargs (`bool`, *optional*, defaults to `False):
366
+ Whether unused keyword arguments of the config are returned.
367
+ return_commit_hash (`bool`, *optional*, defaults to `False):
368
+ Whether the `commit_hash` of the loaded configuration are returned.
369
+
370
+ Returns:
371
+ `dict`:
372
+ A dictionary of all the parameters stored in a JSON configuration file.
373
+
374
+ """
375
+ cache_dir = kwargs.pop("cache_dir", None)
376
+ local_dir = kwargs.pop("local_dir", None)
377
+ local_dir_use_symlinks = kwargs.pop("local_dir_use_symlinks", "auto")
378
+ force_download = kwargs.pop("force_download", False)
379
+ proxies = kwargs.pop("proxies", None)
380
+ token = kwargs.pop("token", None)
381
+ local_files_only = kwargs.pop("local_files_only", False)
382
+ revision = kwargs.pop("revision", None)
383
+ _ = kwargs.pop("mirror", None)
384
+ subfolder = kwargs.pop("subfolder", None)
385
+ user_agent = kwargs.pop("user_agent", {})
386
+ dduf_entries: dict[str, DDUFEntry] | None = kwargs.pop("dduf_entries", None)
387
+
388
+ user_agent = {**user_agent, "file_type": "config"}
389
+ user_agent = http_user_agent(user_agent)
390
+
391
+ pretrained_model_name_or_path = str(pretrained_model_name_or_path)
392
+
393
+ if cls.config_name is None:
394
+ raise ValueError(
395
+ "`self.config_name` is not defined. Note that one should not load a config from "
396
+ "`ConfigMixin`. Please make sure to define `config_name` in a class inheriting from `ConfigMixin`"
397
+ )
398
+ # Custom path for now
399
+ if dduf_entries:
400
+ if subfolder is not None:
401
+ raise ValueError(
402
+ "DDUF file only allow for 1 level of directory (e.g transformer/model1/model.safetentors is not allowed). "
403
+ "Please check the DDUF structure"
404
+ )
405
+ config_file = cls._get_config_file_from_dduf(pretrained_model_name_or_path, dduf_entries)
406
+ elif os.path.isfile(pretrained_model_name_or_path):
407
+ config_file = pretrained_model_name_or_path
408
+ elif os.path.isdir(pretrained_model_name_or_path):
409
+ if subfolder is not None and os.path.isfile(
410
+ os.path.join(pretrained_model_name_or_path, subfolder, cls.config_name)
411
+ ):
412
+ config_file = os.path.join(pretrained_model_name_or_path, subfolder, cls.config_name)
413
+ elif os.path.isfile(os.path.join(pretrained_model_name_or_path, cls.config_name)):
414
+ # Load from a PyTorch checkpoint
415
+ config_file = os.path.join(pretrained_model_name_or_path, cls.config_name)
416
+ else:
417
+ raise EnvironmentError(
418
+ f"Error no file named {cls.config_name} found in directory {pretrained_model_name_or_path}."
419
+ )
420
+ else:
421
+ try:
422
+ # Load from URL or cache if already cached
423
+ config_file = hf_hub_download(
424
+ pretrained_model_name_or_path,
425
+ filename=cls.config_name,
426
+ cache_dir=cache_dir,
427
+ force_download=force_download,
428
+ proxies=proxies,
429
+ local_files_only=local_files_only,
430
+ token=token,
431
+ user_agent=user_agent,
432
+ subfolder=subfolder,
433
+ revision=revision,
434
+ local_dir=local_dir,
435
+ local_dir_use_symlinks=local_dir_use_symlinks,
436
+ )
437
+ except RepositoryNotFoundError:
438
+ raise EnvironmentError(
439
+ f"{pretrained_model_name_or_path} is not a local folder and is not a valid model identifier"
440
+ " listed on 'https://huggingface.co/models'\nIf this is a private repository, make sure to pass a"
441
+ " token having permission to this repo with `token` or log in with `hf auth login`."
442
+ )
443
+ except RevisionNotFoundError:
444
+ raise EnvironmentError(
445
+ f"{revision} is not a valid git identifier (branch name, tag name or commit id) that exists for"
446
+ " this model name. Check the model page at"
447
+ f" 'https://huggingface.co/{pretrained_model_name_or_path}' for available revisions."
448
+ )
449
+ except EntryNotFoundError:
450
+ raise EnvironmentError(
451
+ f"{pretrained_model_name_or_path} does not appear to have a file named {cls.config_name}."
452
+ )
453
+ except HfHubHTTPError as err:
454
+ raise EnvironmentError(
455
+ "There was a specific connection error when trying to load"
456
+ f" {pretrained_model_name_or_path}:\n{err}"
457
+ )
458
+ except ValueError:
459
+ raise EnvironmentError(
460
+ f"We couldn't connect to '{HUGGINGFACE_CO_RESOLVE_ENDPOINT}' to load this model, couldn't find it"
461
+ f" in the cached files and it looks like {pretrained_model_name_or_path} is not the path to a"
462
+ f" directory containing a {cls.config_name} file.\nCheckout your internet connection or see how to"
463
+ " run the library in offline mode at"
464
+ " 'https://huggingface.co/docs/diffusers/installation#offline-mode'."
465
+ )
466
+ except EnvironmentError:
467
+ raise EnvironmentError(
468
+ f"Can't load config for '{pretrained_model_name_or_path}'. If you were trying to load it from "
469
+ "'https://huggingface.co/models', make sure you don't have a local directory with the same name. "
470
+ f"Otherwise, make sure '{pretrained_model_name_or_path}' is the correct path to a directory "
471
+ f"containing a {cls.config_name} file"
472
+ )
473
+ try:
474
+ config_dict = cls._dict_from_json_file(config_file, dduf_entries=dduf_entries)
475
+
476
+ commit_hash = extract_commit_hash(config_file)
477
+ except (json.JSONDecodeError, UnicodeDecodeError):
478
+ raise EnvironmentError(f"It looks like the config file at '{config_file}' is not a valid JSON file.")
479
+
480
+ if not (return_unused_kwargs or return_commit_hash):
481
+ return config_dict
482
+
483
+ outputs = (config_dict,)
484
+
485
+ if return_unused_kwargs:
486
+ outputs += (kwargs,)
487
+
488
+ if return_commit_hash:
489
+ outputs += (commit_hash,)
490
+
491
+ return outputs
492
+
493
+ @staticmethod
494
+ def _get_init_keys(input_class):
495
+ return set(dict(inspect.signature(input_class.__init__).parameters).keys())
496
+
497
+ @classmethod
498
+ def extract_init_dict(cls, config_dict, **kwargs):
499
+ # Skip keys that were not present in the original config, so default __init__ values were used
500
+ used_defaults = config_dict.get("_use_default_values", [])
501
+ config_dict = {k: v for k, v in config_dict.items() if k not in used_defaults and k != "_use_default_values"}
502
+
503
+ # 0. Copy origin config dict
504
+ original_dict = dict(config_dict.items())
505
+
506
+ # 1. Retrieve expected config attributes from __init__ signature
507
+ expected_keys = cls._get_init_keys(cls)
508
+ expected_keys.remove("self")
509
+ # remove general kwargs if present in dict
510
+ if "kwargs" in expected_keys:
511
+ expected_keys.remove("kwargs")
512
+
513
+ # 2. Remove attributes that cannot be expected from expected config attributes
514
+ # remove keys to be ignored
515
+ if len(cls.ignore_for_config) > 0:
516
+ expected_keys = expected_keys - set(cls.ignore_for_config)
517
+
518
+ # load diffusers library to import compatible and original scheduler
519
+ diffusers_library = importlib.import_module(__name__.split(".")[0])
520
+
521
+ if cls.has_compatibles:
522
+ compatible_classes = [c for c in cls._get_compatibles() if not isinstance(c, DummyObject)]
523
+ else:
524
+ compatible_classes = []
525
+
526
+ expected_keys_comp_cls = set()
527
+ for c in compatible_classes:
528
+ expected_keys_c = cls._get_init_keys(c)
529
+ expected_keys_comp_cls = expected_keys_comp_cls.union(expected_keys_c)
530
+ expected_keys_comp_cls = expected_keys_comp_cls - cls._get_init_keys(cls)
531
+ config_dict = {k: v for k, v in config_dict.items() if k not in expected_keys_comp_cls}
532
+
533
+ # remove attributes from orig class that cannot be expected
534
+ orig_cls_name = config_dict.pop("_class_name", cls.__name__)
535
+ if (
536
+ isinstance(orig_cls_name, str)
537
+ and orig_cls_name != cls.__name__
538
+ and hasattr(diffusers_library, orig_cls_name)
539
+ ):
540
+ orig_cls = getattr(diffusers_library, orig_cls_name)
541
+ unexpected_keys_from_orig = cls._get_init_keys(orig_cls) - expected_keys
542
+ config_dict = {k: v for k, v in config_dict.items() if k not in unexpected_keys_from_orig}
543
+ elif not isinstance(orig_cls_name, str) and not isinstance(orig_cls_name, (list, tuple)):
544
+ raise ValueError(
545
+ "Make sure that the `_class_name` is of type string or list of string (for custom pipelines)."
546
+ )
547
+
548
+ # remove private attributes
549
+ config_dict = {k: v for k, v in config_dict.items() if not k.startswith("_")}
550
+
551
+ # remove quantization_config
552
+ config_dict = {k: v for k, v in config_dict.items() if k != "quantization_config"}
553
+
554
+ # 3. Create keyword arguments that will be passed to __init__ from expected keyword arguments
555
+ init_dict = {}
556
+ for key in expected_keys:
557
+ # if config param is passed to kwarg and is present in config dict
558
+ # it should overwrite existing config dict key
559
+ if key in kwargs and key in config_dict:
560
+ config_dict[key] = kwargs.pop(key)
561
+
562
+ if key in kwargs:
563
+ # overwrite key
564
+ init_dict[key] = kwargs.pop(key)
565
+ elif key in config_dict:
566
+ # use value from config dict
567
+ init_dict[key] = config_dict.pop(key)
568
+
569
+ # 4. Give nice warning if unexpected values have been passed
570
+ if len(config_dict) > 0:
571
+ logger.warning(
572
+ f"The config attributes {config_dict} were passed to {cls.__name__}, "
573
+ "but are not expected and will be ignored. Please verify your "
574
+ f"{cls.config_name} configuration file."
575
+ )
576
+
577
+ # 5. Give nice info if config attributes are initialized to default because they have not been passed
578
+ passed_keys = set(init_dict.keys())
579
+ if len(expected_keys - passed_keys) > 0:
580
+ logger.info(
581
+ f"{expected_keys - passed_keys} was not found in config. Values will be initialized to default values."
582
+ )
583
+
584
+ # 6. Define unused keyword arguments
585
+ unused_kwargs = {**config_dict, **kwargs}
586
+
587
+ # 7. Define "hidden" config parameters that were saved for compatible classes
588
+ hidden_config_dict = {k: v for k, v in original_dict.items() if k not in init_dict}
589
+
590
+ return init_dict, unused_kwargs, hidden_config_dict
591
+
592
+ @classmethod
593
+ def _dict_from_json_file(cls, json_file: str | os.PathLike, dduf_entries: dict[str, DDUFEntry] | None = None):
594
+ if dduf_entries:
595
+ text = dduf_entries[json_file].read_text()
596
+ else:
597
+ with open(json_file, "r", encoding="utf-8") as reader:
598
+ text = reader.read()
599
+ return json.loads(text)
600
+
601
+ def __repr__(self):
602
+ return f"{self.__class__.__name__} {self.to_json_string()}"
603
+
604
+ @property
605
+ def config(self) -> dict[str, Any]:
606
+ """
607
+ Returns the config of the class as a frozen dictionary
608
+
609
+ Returns:
610
+ `dict[str, Any]`: Config of the class.
611
+ """
612
+ return self._internal_dict
613
+
614
+ def to_json_string(self) -> str:
615
+ """
616
+ Serializes the configuration instance to a JSON string.
617
+
618
+ Returns:
619
+ `str`:
620
+ String containing all the attributes that make up the configuration instance in JSON format.
621
+ """
622
+ config_dict = self._internal_dict if hasattr(self, "_internal_dict") else {}
623
+ config_dict["_class_name"] = self.__class__.__name__
624
+ config_dict["_diffusers_version"] = __version__
625
+
626
+ def to_json_saveable(value):
627
+ if isinstance(value, np.ndarray):
628
+ value = value.tolist()
629
+ elif isinstance(value, Path):
630
+ value = value.as_posix()
631
+ elif hasattr(value, "to_dict") and callable(value.to_dict):
632
+ value = value.to_dict()
633
+ elif isinstance(value, list):
634
+ value = [to_json_saveable(v) for v in value]
635
+ return value
636
+
637
+ if "quantization_config" in config_dict:
638
+ config_dict["quantization_config"] = (
639
+ config_dict.quantization_config.to_dict()
640
+ if not isinstance(config_dict.quantization_config, dict)
641
+ else config_dict.quantization_config
642
+ )
643
+
644
+ config_dict = {k: to_json_saveable(v) for k, v in config_dict.items()}
645
+ # Don't save "_ignore_files" or "_use_default_values"
646
+ config_dict.pop("_ignore_files", None)
647
+ config_dict.pop("_use_default_values", None)
648
+ # pop the `_pre_quantization_dtype` as torch.dtypes are not serializable.
649
+ _ = config_dict.pop("_pre_quantization_dtype", None)
650
+
651
+ if getattr(self, "_auto_class", None) is not None:
652
+ module = self.__class__.__module__.split(".")[-1]
653
+ auto_map = config_dict.get("auto_map", {})
654
+ auto_map[self._auto_class] = f"{module}.{self.__class__.__name__}"
655
+ config_dict["auto_map"] = auto_map
656
+
657
+ return json.dumps(config_dict, indent=2, sort_keys=True) + "\n"
658
+
659
+ def to_json_file(self, json_file_path: str | os.PathLike):
660
+ """
661
+ Save the configuration instance's parameters to a JSON file.
662
+
663
+ Args:
664
+ json_file_path (`str` or `os.PathLike`):
665
+ Path to the JSON file to save a configuration instance's parameters.
666
+ """
667
+ with open(json_file_path, "w", encoding="utf-8") as writer:
668
+ writer.write(self.to_json_string())
669
+
670
+ @classmethod
671
+ def _get_config_file_from_dduf(cls, pretrained_model_name_or_path: str, dduf_entries: dict[str, DDUFEntry]):
672
+ # paths inside a DDUF file must always be "/"
673
+ config_file = (
674
+ cls.config_name
675
+ if pretrained_model_name_or_path == ""
676
+ else "/".join([pretrained_model_name_or_path, cls.config_name])
677
+ )
678
+ if config_file not in dduf_entries:
679
+ raise ValueError(
680
+ f"We did not manage to find the file {config_file} in the dduf file. We only have the following files {dduf_entries.keys()}"
681
+ )
682
+ return config_file
683
+
684
+
685
+ def register_to_config(init):
686
+ r"""
687
+ Decorator to apply on the init of classes inheriting from [`ConfigMixin`] so that all the arguments are
688
+ automatically sent to `self.register_for_config`. To ignore a specific argument accepted by the init but that
689
+ shouldn't be registered in the config, use the `ignore_for_config` class variable
690
+
691
+ Warning: Once decorated, all private arguments (beginning with an underscore) are trashed and not sent to the init!
692
+ """
693
+
694
+ @functools.wraps(init)
695
+ def inner_init(self, *args, **kwargs):
696
+ # Ignore private kwargs in the init.
697
+ init_kwargs = {k: v for k, v in kwargs.items() if not k.startswith("_")}
698
+ config_init_kwargs = {k: v for k, v in kwargs.items() if k.startswith("_")}
699
+ if not isinstance(self, ConfigMixin):
700
+ raise RuntimeError(
701
+ f"`@register_for_config` was applied to {self.__class__.__name__} init method, but this class does "
702
+ "not inherit from `ConfigMixin`."
703
+ )
704
+
705
+ ignore = getattr(self, "ignore_for_config", [])
706
+ # Get positional arguments aligned with kwargs
707
+ new_kwargs = {}
708
+ signature = inspect.signature(init)
709
+ parameters = {
710
+ name: p.default for i, (name, p) in enumerate(signature.parameters.items()) if i > 0 and name not in ignore
711
+ }
712
+ for arg, name in zip(args, parameters.keys()):
713
+ new_kwargs[name] = arg
714
+
715
+ # Then add all kwargs
716
+ new_kwargs.update(
717
+ {
718
+ k: init_kwargs.get(k, default)
719
+ for k, default in parameters.items()
720
+ if k not in ignore and k not in new_kwargs
721
+ }
722
+ )
723
+
724
+ # Take note of the parameters that were not present in the loaded config
725
+ if len(set(new_kwargs.keys()) - set(init_kwargs)) > 0:
726
+ new_kwargs["_use_default_values"] = list(set(new_kwargs.keys()) - set(init_kwargs))
727
+
728
+ new_kwargs = {**config_init_kwargs, **new_kwargs}
729
+ getattr(self, "register_to_config")(**new_kwargs)
730
+ init(self, *args, **init_kwargs)
731
+
732
+ return inner_init
733
+
734
+
735
+ class LegacyConfigMixin(ConfigMixin):
736
+ r"""
737
+ A subclass of `ConfigMixin` to resolve class mapping from legacy classes (like `Transformer2DModel`) to more
738
+ pipeline-specific classes (like `DiTTransformer2DModel`).
739
+ """
740
+
741
+ @classmethod
742
+ def from_config(cls, config: FrozenDict | dict[str, Any] = None, return_unused_kwargs=False, **kwargs):
743
+ # To prevent dependency import problem.
744
+ from .models.model_loading_utils import _fetch_remapped_cls_from_config
745
+
746
+ # resolve remapping
747
+ remapped_class = _fetch_remapped_cls_from_config(config, cls)
748
+
749
+ if remapped_class is cls:
750
+ return super(LegacyConfigMixin, remapped_class).from_config(config, return_unused_kwargs, **kwargs)
751
+ else:
752
+ return remapped_class.from_config(config, return_unused_kwargs, **kwargs)
diffusers/dependency_versions_check.py ADDED
@@ -0,0 +1,34 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 The HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ from .dependency_versions_table import deps
16
+ from .utils.versions import require_version, require_version_core
17
+
18
+
19
+ # define which module versions we always want to check at run time
20
+ # (usually the ones defined in `install_requires` in setup.py)
21
+ #
22
+ # order specific notes:
23
+ # - tqdm must be checked before tokenizers
24
+
25
+ pkgs_to_check_at_runtime = "python requests filelock numpy".split()
26
+ for pkg in pkgs_to_check_at_runtime:
27
+ if pkg in deps:
28
+ require_version_core(deps[pkg])
29
+ else:
30
+ raise ValueError(f"can't find {pkg} in {deps.keys()}, check dependency_versions_table.py")
31
+
32
+
33
+ def dep_version_check(pkg, hint=None):
34
+ require_version(deps[pkg], hint)
diffusers/dependency_versions_table.py ADDED
@@ -0,0 +1,57 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # THIS FILE HAS BEEN AUTOGENERATED. To update:
2
+ # 1. modify the `_deps` dict in setup.py
3
+ # 2. run `make deps_table_update`
4
+ deps = {
5
+ "Pillow": "Pillow",
6
+ "accelerate": "accelerate>=0.31.0",
7
+ "datasets": "datasets",
8
+ "filelock": "filelock",
9
+ "ftfy": "ftfy",
10
+ "hf-doc-builder": "hf-doc-builder>=0.3.0",
11
+ "httpx": "httpx<1.0.0",
12
+ "huggingface-hub": "huggingface-hub>=1.23.0,<2.0",
13
+ "requests-mock": "requests-mock==1.10.0",
14
+ "importlib_metadata": "importlib_metadata",
15
+ "invisible-watermark": "invisible-watermark>=0.2.0",
16
+ "isort": "isort>=5.5.4",
17
+ "Jinja2": "Jinja2",
18
+ "torchsde": "torchsde",
19
+ "note_seq": "note_seq",
20
+ "librosa": "librosa",
21
+ "llvmlite": "llvmlite>=0.40.0",
22
+ "numba": "numba>=0.57.0",
23
+ "numpy": "numpy",
24
+ "parameterized": "parameterized",
25
+ "peft": "peft>=0.17.0",
26
+ "protobuf": "protobuf>=3.20.3,<4",
27
+ "pytest": "pytest",
28
+ "pytest-timeout": "pytest-timeout",
29
+ "pytest-xdist": "pytest-xdist",
30
+ "python": "python>=3.10.0",
31
+ "ruff": "ruff==0.9.10",
32
+ "safetensors": "safetensors>=0.8.0",
33
+ "sentencepiece": "sentencepiece>=0.1.91,!=0.1.92",
34
+ "GitPython": "GitPython<3.1.19",
35
+ "scipy": "scipy",
36
+ "onnx": "onnx",
37
+ "optimum_quanto": "optimum_quanto>=0.2.6",
38
+ "gguf": "gguf>=0.10.0",
39
+ "auto-round": "auto-round>=0.13.0",
40
+ "torchao": "torchao>=0.7.0",
41
+ "bitsandbytes": "bitsandbytes>=0.43.3",
42
+ "nvidia_modelopt[hf]": "nvidia_modelopt[hf]>=0.33.1",
43
+ "sdnq": "sdnq>=0.2.2",
44
+ "regex": "regex!=2019.12.17",
45
+ "requests": "requests",
46
+ "tensorboard": "tensorboard",
47
+ "tiktoken": "tiktoken>=0.7.0",
48
+ "torch": "torch>=2.6",
49
+ "torchvision": "torchvision",
50
+ "transformers": "transformers>=4.41.2",
51
+ "urllib3": "urllib3<=2.0.0",
52
+ "black": "black",
53
+ "phonemizer": "phonemizer",
54
+ "opencv-python": "opencv-python",
55
+ "timm": "timm",
56
+ "flashpack": "flashpack",
57
+ }
diffusers/experimental/README.md ADDED
@@ -0,0 +1,5 @@
 
 
 
 
 
 
1
+ # 🧨 Diffusers Experimental
2
+
3
+ We are adding experimental code to support novel applications and usages of the Diffusers library.
4
+ Currently, the following experiments are supported:
5
+ * Reinforcement learning via an implementation of the [Diffuser](https://huggingface.co/papers/2205.09991) model.
diffusers/experimental/__init__.py ADDED
@@ -0,0 +1 @@
 
 
1
+ from .rl import ValueGuidedRLPipeline
diffusers/experimental/rl/__init__.py ADDED
@@ -0,0 +1 @@
 
 
1
+ from .value_guided_sampling import ValueGuidedRLPipeline
diffusers/experimental/rl/value_guided_sampling.py ADDED
@@ -0,0 +1,153 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 The HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ import numpy as np
16
+ import torch
17
+ import tqdm
18
+
19
+ from ...models.unets.unet_1d import UNet1DModel
20
+ from ...pipelines import DiffusionPipeline
21
+ from ...utils.dummy_pt_objects import DDPMScheduler
22
+ from ...utils.torch_utils import randn_tensor
23
+
24
+
25
+ class ValueGuidedRLPipeline(DiffusionPipeline):
26
+ r"""
27
+ Pipeline for value-guided sampling from a diffusion model trained to predict sequences of states.
28
+
29
+ This model inherits from [`DiffusionPipeline`]. Check the superclass documentation for the generic methods
30
+ implemented for all pipelines (downloading, saving, running on a particular device, etc.).
31
+
32
+ Parameters:
33
+ value_function ([`UNet1DModel`]):
34
+ A specialized UNet for fine-tuning trajectories base on reward.
35
+ unet ([`UNet1DModel`]):
36
+ UNet architecture to denoise the encoded trajectories.
37
+ scheduler ([`SchedulerMixin`]):
38
+ A scheduler to be used in combination with `unet` to denoise the encoded trajectories. Default for this
39
+ application is [`DDPMScheduler`].
40
+ env ():
41
+ An environment following the OpenAI gym API to act in. For now only Hopper has pretrained models.
42
+ """
43
+
44
+ def __init__(
45
+ self,
46
+ value_function: UNet1DModel,
47
+ unet: UNet1DModel,
48
+ scheduler: DDPMScheduler,
49
+ env,
50
+ ):
51
+ super().__init__()
52
+
53
+ self.register_modules(value_function=value_function, unet=unet, scheduler=scheduler, env=env)
54
+
55
+ self.data = env.get_dataset()
56
+ self.means = {}
57
+ for key in self.data.keys():
58
+ try:
59
+ self.means[key] = self.data[key].mean()
60
+ except: # noqa: E722
61
+ pass
62
+ self.stds = {}
63
+ for key in self.data.keys():
64
+ try:
65
+ self.stds[key] = self.data[key].std()
66
+ except: # noqa: E722
67
+ pass
68
+ self.state_dim = env.observation_space.shape[0]
69
+ self.action_dim = env.action_space.shape[0]
70
+
71
+ def normalize(self, x_in, key):
72
+ return (x_in - self.means[key]) / self.stds[key]
73
+
74
+ def de_normalize(self, x_in, key):
75
+ return x_in * self.stds[key] + self.means[key]
76
+
77
+ def to_torch(self, x_in):
78
+ if isinstance(x_in, dict):
79
+ return {k: self.to_torch(v) for k, v in x_in.items()}
80
+ elif torch.is_tensor(x_in):
81
+ return x_in.to(self.unet.device)
82
+ return torch.tensor(x_in, device=self.unet.device)
83
+
84
+ def reset_x0(self, x_in, cond, act_dim):
85
+ for key, val in cond.items():
86
+ x_in[:, key, act_dim:] = val.clone()
87
+ return x_in
88
+
89
+ def run_diffusion(self, x, conditions, n_guide_steps, scale):
90
+ batch_size = x.shape[0]
91
+ y = None
92
+ for i in tqdm.tqdm(self.scheduler.timesteps):
93
+ # create batch of timesteps to pass into model
94
+ timesteps = torch.full((batch_size,), i, device=self.unet.device, dtype=torch.long)
95
+ for _ in range(n_guide_steps):
96
+ with torch.enable_grad():
97
+ x.requires_grad_()
98
+
99
+ # permute to match dimension for pre-trained models
100
+ y = self.value_function(x.permute(0, 2, 1), timesteps).sample
101
+ grad = torch.autograd.grad([y.sum()], [x])[0]
102
+
103
+ posterior_variance = self.scheduler._get_variance(i)
104
+ model_std = torch.exp(0.5 * posterior_variance)
105
+ grad = model_std * grad
106
+
107
+ grad[timesteps < 2] = 0
108
+ x = x.detach()
109
+ x = x + scale * grad
110
+ x = self.reset_x0(x, conditions, self.action_dim)
111
+
112
+ prev_x = self.unet(x.permute(0, 2, 1), timesteps).sample.permute(0, 2, 1)
113
+
114
+ # TODO: verify deprecation of this kwarg
115
+ x = self.scheduler.step(prev_x, i, x)["prev_sample"]
116
+
117
+ # apply conditions to the trajectory (set the initial state)
118
+ x = self.reset_x0(x, conditions, self.action_dim)
119
+ x = self.to_torch(x)
120
+ return x, y
121
+
122
+ def __call__(self, obs, batch_size=64, planning_horizon=32, n_guide_steps=2, scale=0.1):
123
+ # normalize the observations and create batch dimension
124
+ obs = self.normalize(obs, "observations")
125
+ obs = obs[None].repeat(batch_size, axis=0)
126
+
127
+ conditions = {0: self.to_torch(obs)}
128
+ shape = (batch_size, planning_horizon, self.state_dim + self.action_dim)
129
+
130
+ # generate initial noise and apply our conditions (to make the trajectories start at current state)
131
+ x1 = randn_tensor(shape, device=self.unet.device)
132
+ x = self.reset_x0(x1, conditions, self.action_dim)
133
+ x = self.to_torch(x)
134
+
135
+ # run the diffusion process
136
+ x, y = self.run_diffusion(x, conditions, n_guide_steps, scale)
137
+
138
+ # sort output trajectories by value
139
+ sorted_idx = y.argsort(0, descending=True).squeeze()
140
+ sorted_values = x[sorted_idx]
141
+ actions = sorted_values[:, :, : self.action_dim]
142
+ actions = actions.detach().cpu().numpy()
143
+ denorm_actions = self.de_normalize(actions, key="actions")
144
+
145
+ # select the action with the highest value
146
+ if y is not None:
147
+ selected_index = 0
148
+ else:
149
+ # if we didn't run value guiding, select a random action
150
+ selected_index = np.random.randint(0, batch_size)
151
+
152
+ denorm_actions = denorm_actions[selected_index, 0]
153
+ return denorm_actions
diffusers/guiders/__init__.py ADDED
@@ -0,0 +1,31 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 The HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+
16
+ from ..utils import is_torch_available, logging
17
+
18
+
19
+ if is_torch_available():
20
+ from .adaptive_projected_guidance import AdaptiveProjectedGuidance
21
+ from .adaptive_projected_guidance_mix import AdaptiveProjectedMixGuidance
22
+ from .auto_guidance import AutoGuidance
23
+ from .classifier_free_guidance import ClassifierFreeGuidance
24
+ from .classifier_free_zero_star_guidance import ClassifierFreeZeroStarGuidance
25
+ from .frequency_decoupled_guidance import FrequencyDecoupledGuidance
26
+ from .guider_utils import BaseGuidance
27
+ from .magnitude_aware_guidance import MagnitudeAwareGuidance
28
+ from .perturbed_attention_guidance import PerturbedAttentionGuidance
29
+ from .skip_layer_guidance import SkipLayerGuidance
30
+ from .smoothed_energy_guidance import SmoothedEnergyGuidance
31
+ from .tangential_classifier_free_guidance import TangentialClassifierFreeGuidance
diffusers/guiders/adaptive_projected_guidance.py ADDED
@@ -0,0 +1,253 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 The HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ from __future__ import annotations
16
+
17
+ import math
18
+ from typing import TYPE_CHECKING
19
+
20
+ import torch
21
+
22
+ from ..configuration_utils import register_to_config
23
+ from .guider_utils import BaseGuidance, GuiderOutput, rescale_noise_cfg
24
+
25
+
26
+ if TYPE_CHECKING:
27
+ from ..modular_pipelines.modular_pipeline import BlockState
28
+
29
+
30
+ class AdaptiveProjectedGuidance(BaseGuidance):
31
+ """
32
+ Adaptive Projected Guidance (APG): https://huggingface.co/papers/2410.02416
33
+
34
+ Args:
35
+ guidance_scale (`float`, defaults to `7.5`):
36
+ The scale parameter for classifier-free guidance. Higher values result in stronger conditioning on the text
37
+ prompt, while lower values allow for more freedom in generation. Higher values may lead to saturation and
38
+ deterioration of image quality.
39
+ adaptive_projected_guidance_momentum (`float`, defaults to `None`):
40
+ The momentum parameter for the adaptive projected guidance. Disabled if set to `None`.
41
+ adaptive_projected_guidance_rescale (`float`, defaults to `15.0`):
42
+ The rescale factor applied to the noise predictions. This is used to improve image quality and fix
43
+ adaptive_projected_guidance_norm_dim (`int` or `tuple[int]`, *optional*):
44
+ Dimension(s) over which to compute the APG norm and projection. If omitted, all non-batch dimensions are
45
+ used, preserving the original behavior.
46
+ guidance_rescale (`float`, defaults to `0.0`):
47
+ The rescale factor applied to the noise predictions. This is used to improve image quality and fix
48
+ overexposure. Based on Section 3.4 from [Common Diffusion Noise Schedules and Sample Steps are
49
+ Flawed](https://huggingface.co/papers/2305.08891).
50
+ use_original_formulation (`bool`, defaults to `False`):
51
+ Whether to use the original formulation of classifier-free guidance as proposed in the paper. By default,
52
+ we use the diffusers-native implementation that has been in the codebase for a long time. See
53
+ [~guiders.classifier_free_guidance.ClassifierFreeGuidance] for more details.
54
+ start (`float`, defaults to `0.0`):
55
+ The fraction of the total number of denoising steps after which guidance starts.
56
+ stop (`float`, defaults to `1.0`):
57
+ The fraction of the total number of denoising steps after which guidance stops.
58
+ """
59
+
60
+ _input_predictions = ["pred_cond", "pred_uncond"]
61
+
62
+ @register_to_config
63
+ def __init__(
64
+ self,
65
+ guidance_scale: float = 7.5,
66
+ adaptive_projected_guidance_momentum: float | None = None,
67
+ adaptive_projected_guidance_rescale: float = 15.0,
68
+ adaptive_projected_guidance_norm_dim: int | tuple[int, ...] | None = None,
69
+ eta: float = 1.0,
70
+ guidance_rescale: float = 0.0,
71
+ use_original_formulation: bool = False,
72
+ start: float = 0.0,
73
+ stop: float = 1.0,
74
+ enabled: bool = True,
75
+ ):
76
+ super().__init__(start, stop, enabled)
77
+
78
+ self.guidance_scale = guidance_scale
79
+ self.adaptive_projected_guidance_momentum = adaptive_projected_guidance_momentum
80
+ self.adaptive_projected_guidance_rescale = adaptive_projected_guidance_rescale
81
+ self.adaptive_projected_guidance_norm_dim = adaptive_projected_guidance_norm_dim
82
+ self.eta = eta
83
+ self.guidance_rescale = guidance_rescale
84
+ self.use_original_formulation = use_original_formulation
85
+ self.momentum_buffer = None
86
+
87
+ def prepare_inputs(self, data: dict[str, tuple[torch.Tensor, torch.Tensor]]) -> list["BlockState"]:
88
+ if self._step == 0:
89
+ if self.adaptive_projected_guidance_momentum is not None:
90
+ self.momentum_buffer = MomentumBuffer(self.adaptive_projected_guidance_momentum)
91
+ tuple_indices = [0] if self.num_conditions == 1 else [0, 1]
92
+ data_batches = []
93
+ for tuple_idx, input_prediction in zip(tuple_indices, self._input_predictions):
94
+ data_batch = self._prepare_batch(data, tuple_idx, input_prediction)
95
+ data_batches.append(data_batch)
96
+ return data_batches
97
+
98
+ def prepare_inputs_from_block_state(
99
+ self, data: "BlockState", input_fields: dict[str, str | tuple[str, str]]
100
+ ) -> list["BlockState"]:
101
+ if self._step == 0:
102
+ if self.adaptive_projected_guidance_momentum is not None:
103
+ self.momentum_buffer = MomentumBuffer(self.adaptive_projected_guidance_momentum)
104
+ tuple_indices = [0] if self.num_conditions == 1 else [0, 1]
105
+ data_batches = []
106
+ for tuple_idx, input_prediction in zip(tuple_indices, self._input_predictions):
107
+ data_batch = self._prepare_batch_from_block_state(input_fields, data, tuple_idx, input_prediction)
108
+ data_batches.append(data_batch)
109
+ return data_batches
110
+
111
+ def forward(self, pred_cond: torch.Tensor, pred_uncond: torch.Tensor | None = None) -> GuiderOutput:
112
+ pred = None
113
+
114
+ if not self._is_apg_enabled():
115
+ pred = pred_cond
116
+ else:
117
+ pred = normalized_guidance(
118
+ pred_cond,
119
+ pred_uncond,
120
+ self.guidance_scale,
121
+ self.momentum_buffer,
122
+ self.eta,
123
+ self.adaptive_projected_guidance_rescale,
124
+ self.use_original_formulation,
125
+ self.adaptive_projected_guidance_norm_dim,
126
+ )
127
+
128
+ if self.guidance_rescale > 0.0:
129
+ pred = rescale_noise_cfg(pred, pred_cond, self.guidance_rescale)
130
+
131
+ return GuiderOutput(pred=pred, pred_cond=pred_cond, pred_uncond=pred_uncond)
132
+
133
+ @property
134
+ def is_conditional(self) -> bool:
135
+ return self._count_prepared == 1
136
+
137
+ @property
138
+ def num_conditions(self) -> int:
139
+ num_conditions = 1
140
+ if self._is_apg_enabled():
141
+ num_conditions += 1
142
+ return num_conditions
143
+
144
+ def _is_apg_enabled(self) -> bool:
145
+ if not self._enabled:
146
+ return False
147
+
148
+ is_within_range = True
149
+ if self._num_inference_steps is not None:
150
+ skip_start_step = int(self._start * self._num_inference_steps)
151
+ skip_stop_step = int(self._stop * self._num_inference_steps)
152
+ is_within_range = skip_start_step <= self._step < skip_stop_step
153
+
154
+ is_close = False
155
+ if self.use_original_formulation:
156
+ is_close = math.isclose(self.guidance_scale, 0.0)
157
+ else:
158
+ is_close = math.isclose(self.guidance_scale, 1.0)
159
+
160
+ return is_within_range and not is_close
161
+
162
+
163
+ class MomentumBuffer:
164
+ def __init__(self, momentum: float):
165
+ self.momentum = momentum
166
+ self.running_average = 0
167
+
168
+ def update(self, update_value: torch.Tensor):
169
+ new_average = self.momentum * self.running_average
170
+ self.running_average = update_value + new_average
171
+
172
+ def __repr__(self) -> str:
173
+ """
174
+ Returns a string representation showing momentum, shape, statistics, and a slice of the running_average.
175
+ """
176
+ if isinstance(self.running_average, torch.Tensor):
177
+ shape = tuple(self.running_average.shape)
178
+
179
+ # Calculate statistics
180
+ with torch.no_grad():
181
+ stats = {
182
+ "mean": self.running_average.mean().item(),
183
+ "std": self.running_average.std().item(),
184
+ "min": self.running_average.min().item(),
185
+ "max": self.running_average.max().item(),
186
+ }
187
+
188
+ # Get a slice (max 3 elements per dimension)
189
+ slice_indices = tuple(slice(None, min(3, dim)) for dim in shape)
190
+ sliced_data = self.running_average[slice_indices]
191
+
192
+ # Format the slice for display (convert to float32 for numpy compatibility with bfloat16)
193
+ slice_str = str(sliced_data.detach().float().cpu().numpy())
194
+ if len(slice_str) > 200: # Truncate if too long
195
+ slice_str = slice_str[:200] + "..."
196
+
197
+ stats_str = ", ".join([f"{k}={v:.4f}" for k, v in stats.items()])
198
+
199
+ return (
200
+ f"MomentumBuffer(\n"
201
+ f" momentum={self.momentum},\n"
202
+ f" shape={shape},\n"
203
+ f" stats=[{stats_str}],\n"
204
+ f" slice={slice_str}\n"
205
+ f")"
206
+ )
207
+ else:
208
+ return f"MomentumBuffer(momentum={self.momentum}, running_average={self.running_average})"
209
+
210
+
211
+ def normalized_guidance(
212
+ pred_cond: torch.Tensor,
213
+ pred_uncond: torch.Tensor,
214
+ guidance_scale: float,
215
+ momentum_buffer: MomentumBuffer | None = None,
216
+ eta: float = 1.0,
217
+ norm_threshold: float = 0.0,
218
+ use_original_formulation: bool = False,
219
+ norm_dim: int | tuple[int, ...] | None = None,
220
+ ):
221
+ diff = pred_cond - pred_uncond
222
+ if norm_dim is None:
223
+ dim = [-i for i in range(1, len(diff.shape))]
224
+ elif isinstance(norm_dim, int):
225
+ dim = [norm_dim]
226
+ else:
227
+ dim = list(norm_dim)
228
+
229
+ if momentum_buffer is not None:
230
+ momentum_buffer.update(diff)
231
+ diff = momentum_buffer.running_average
232
+
233
+ if norm_threshold > 0:
234
+ ones = torch.ones_like(diff)
235
+ diff_norm = diff.norm(p=2, dim=dim, keepdim=True)
236
+ scale_factor = torch.minimum(ones, norm_threshold / diff_norm)
237
+ diff = diff * scale_factor
238
+
239
+ if diff.device.type in {"mps", "npu"}:
240
+ v0, v1 = diff.cpu().double(), pred_cond.cpu().double()
241
+ else:
242
+ v0, v1 = diff.double(), pred_cond.double()
243
+ v1 = torch.nn.functional.normalize(v1, dim=dim)
244
+ v0_parallel = (v0 * v1).sum(dim=dim, keepdim=True) * v1
245
+ v0_orthogonal = v0 - v0_parallel
246
+ diff_parallel = v0_parallel.to(device=diff.device, dtype=diff.dtype)
247
+ diff_orthogonal = v0_orthogonal.to(device=diff.device, dtype=diff.dtype)
248
+ normalized_update = diff_orthogonal + eta * diff_parallel
249
+
250
+ pred = pred_cond if use_original_formulation else pred_uncond
251
+ pred = pred + guidance_scale * normalized_update
252
+
253
+ return pred
diffusers/guiders/adaptive_projected_guidance_mix.py ADDED
@@ -0,0 +1,297 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 The HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ import math
16
+ from typing import TYPE_CHECKING
17
+
18
+ import torch
19
+
20
+ from ..configuration_utils import register_to_config
21
+ from .guider_utils import BaseGuidance, GuiderOutput, rescale_noise_cfg
22
+
23
+
24
+ if TYPE_CHECKING:
25
+ from ..modular_pipelines.modular_pipeline import BlockState
26
+
27
+
28
+ class AdaptiveProjectedMixGuidance(BaseGuidance):
29
+ """
30
+ Adaptive Projected Guidance (APG) https://huggingface.co/papers/2410.02416 combined with Classifier-Free Guidance
31
+ (CFG). This guider is used in HunyuanImage2.1 https://github.com/Tencent-Hunyuan/HunyuanImage-2.1
32
+
33
+ Args:
34
+ guidance_scale (`float`, defaults to `7.5`):
35
+ The scale parameter for classifier-free guidance. Higher values result in stronger conditioning on the text
36
+ prompt, while lower values allow for more freedom in generation. Higher values may lead to saturation and
37
+ deterioration of image quality.
38
+ adaptive_projected_guidance_momentum (`float`, defaults to `None`):
39
+ The momentum parameter for the adaptive projected guidance. Disabled if set to `None`.
40
+ adaptive_projected_guidance_rescale (`float`, defaults to `15.0`):
41
+ The rescale factor applied to the noise predictions for adaptive projected guidance. This is used to
42
+ improve image quality and fix
43
+ guidance_rescale (`float`, defaults to `0.0`):
44
+ The rescale factor applied to the noise predictions for classifier-free guidance. This is used to improve
45
+ image quality and fix overexposure. Based on Section 3.4 from [Common Diffusion Noise Schedules and Sample
46
+ Steps are Flawed](https://huggingface.co/papers/2305.08891).
47
+ use_original_formulation (`bool`, defaults to `False`):
48
+ Whether to use the original formulation of classifier-free guidance as proposed in the paper. By default,
49
+ we use the diffusers-native implementation that has been in the codebase for a long time. See
50
+ [~guiders.classifier_free_guidance.ClassifierFreeGuidance] for more details.
51
+ start (`float`, defaults to `0.0`):
52
+ The fraction of the total number of denoising steps after which the classifier-free guidance starts.
53
+ stop (`float`, defaults to `1.0`):
54
+ The fraction of the total number of denoising steps after which the classifier-free guidance stops.
55
+ adaptive_projected_guidance_start_step (`int`, defaults to `5`):
56
+ The step at which the adaptive projected guidance starts (before this step, classifier-free guidance is
57
+ used, and momentum buffer is updated).
58
+ enabled (`bool`, defaults to `True`):
59
+ Whether this guidance is enabled.
60
+ """
61
+
62
+ _input_predictions = ["pred_cond", "pred_uncond"]
63
+
64
+ @register_to_config
65
+ def __init__(
66
+ self,
67
+ guidance_scale: float = 3.5,
68
+ guidance_rescale: float = 0.0,
69
+ adaptive_projected_guidance_scale: float = 10.0,
70
+ adaptive_projected_guidance_momentum: float = -0.5,
71
+ adaptive_projected_guidance_rescale: float = 10.0,
72
+ eta: float = 0.0,
73
+ use_original_formulation: bool = False,
74
+ start: float = 0.0,
75
+ stop: float = 1.0,
76
+ adaptive_projected_guidance_start_step: int = 5,
77
+ enabled: bool = True,
78
+ ):
79
+ super().__init__(start, stop, enabled)
80
+
81
+ self.guidance_scale = guidance_scale
82
+ self.guidance_rescale = guidance_rescale
83
+ self.adaptive_projected_guidance_scale = adaptive_projected_guidance_scale
84
+ self.adaptive_projected_guidance_momentum = adaptive_projected_guidance_momentum
85
+ self.adaptive_projected_guidance_rescale = adaptive_projected_guidance_rescale
86
+ self.eta = eta
87
+ self.adaptive_projected_guidance_start_step = adaptive_projected_guidance_start_step
88
+ self.use_original_formulation = use_original_formulation
89
+ self.momentum_buffer = None
90
+
91
+ def prepare_inputs(self, data: dict[str, tuple[torch.Tensor, torch.Tensor]]) -> list["BlockState"]:
92
+ if self._step == 0:
93
+ if self.adaptive_projected_guidance_momentum is not None:
94
+ self.momentum_buffer = MomentumBuffer(self.adaptive_projected_guidance_momentum)
95
+ tuple_indices = [0] if self.num_conditions == 1 else [0, 1]
96
+ data_batches = []
97
+ for tuple_idx, input_prediction in zip(tuple_indices, self._input_predictions):
98
+ data_batch = self._prepare_batch(data, tuple_idx, input_prediction)
99
+ data_batches.append(data_batch)
100
+ return data_batches
101
+
102
+ def prepare_inputs_from_block_state(
103
+ self, data: "BlockState", input_fields: dict[str, str | tuple[str, str]]
104
+ ) -> list["BlockState"]:
105
+ if self._step == 0:
106
+ if self.adaptive_projected_guidance_momentum is not None:
107
+ self.momentum_buffer = MomentumBuffer(self.adaptive_projected_guidance_momentum)
108
+ tuple_indices = [0] if self.num_conditions == 1 else [0, 1]
109
+ data_batches = []
110
+ for tuple_idx, input_prediction in zip(tuple_indices, self._input_predictions):
111
+ data_batch = self._prepare_batch_from_block_state(input_fields, data, tuple_idx, input_prediction)
112
+ data_batches.append(data_batch)
113
+ return data_batches
114
+
115
+ def forward(self, pred_cond: torch.Tensor, pred_uncond: torch.Tensor | None = None) -> GuiderOutput:
116
+ pred = None
117
+
118
+ # no guidance
119
+ if not self._is_cfg_enabled():
120
+ pred = pred_cond
121
+
122
+ # CFG + update momentum buffer
123
+ elif not self._is_apg_enabled():
124
+ if self.momentum_buffer is not None:
125
+ update_momentum_buffer(pred_cond, pred_uncond, self.momentum_buffer)
126
+ # CFG + update momentum buffer
127
+ shift = pred_cond - pred_uncond
128
+ pred = pred_cond if self.use_original_formulation else pred_uncond
129
+ pred = pred + self.guidance_scale * shift
130
+
131
+ # APG
132
+ elif self._is_apg_enabled():
133
+ pred = normalized_guidance(
134
+ pred_cond,
135
+ pred_uncond,
136
+ self.adaptive_projected_guidance_scale,
137
+ self.momentum_buffer,
138
+ self.eta,
139
+ self.adaptive_projected_guidance_rescale,
140
+ self.use_original_formulation,
141
+ )
142
+
143
+ if self.guidance_rescale > 0.0:
144
+ pred = rescale_noise_cfg(pred, pred_cond, self.guidance_rescale)
145
+
146
+ return GuiderOutput(pred=pred, pred_cond=pred_cond, pred_uncond=pred_uncond)
147
+
148
+ @property
149
+ def is_conditional(self) -> bool:
150
+ return self._count_prepared == 1
151
+
152
+ @property
153
+ def num_conditions(self) -> int:
154
+ num_conditions = 1
155
+ if self._is_apg_enabled() or self._is_cfg_enabled():
156
+ num_conditions += 1
157
+ return num_conditions
158
+
159
+ # Copied from diffusers.guiders.classifier_free_guidance.ClassifierFreeGuidance._is_cfg_enabled
160
+ def _is_cfg_enabled(self) -> bool:
161
+ if not self._enabled:
162
+ return False
163
+
164
+ is_within_range = True
165
+ if self._num_inference_steps is not None:
166
+ skip_start_step = int(self._start * self._num_inference_steps)
167
+ skip_stop_step = int(self._stop * self._num_inference_steps)
168
+ is_within_range = skip_start_step <= self._step < skip_stop_step
169
+
170
+ is_close = False
171
+ if self.use_original_formulation:
172
+ is_close = math.isclose(self.guidance_scale, 0.0)
173
+ else:
174
+ is_close = math.isclose(self.guidance_scale, 1.0)
175
+
176
+ return is_within_range and not is_close
177
+
178
+ def _is_apg_enabled(self) -> bool:
179
+ if not self._enabled:
180
+ return False
181
+
182
+ if not self._is_cfg_enabled():
183
+ return False
184
+
185
+ is_within_range = False
186
+ if self._step is not None:
187
+ is_within_range = self._step > self.adaptive_projected_guidance_start_step
188
+
189
+ is_close = False
190
+ if self.use_original_formulation:
191
+ is_close = math.isclose(self.adaptive_projected_guidance_scale, 0.0)
192
+ else:
193
+ is_close = math.isclose(self.adaptive_projected_guidance_scale, 1.0)
194
+
195
+ return is_within_range and not is_close
196
+
197
+ def get_state(self):
198
+ state = super().get_state()
199
+ state["momentum_buffer"] = self.momentum_buffer
200
+ state["is_apg_enabled"] = self._is_apg_enabled()
201
+ state["is_cfg_enabled"] = self._is_cfg_enabled()
202
+ return state
203
+
204
+
205
+ # Copied from diffusers.guiders.adaptive_projected_guidance.MomentumBuffer
206
+ class MomentumBuffer:
207
+ def __init__(self, momentum: float):
208
+ self.momentum = momentum
209
+ self.running_average = 0
210
+
211
+ def update(self, update_value: torch.Tensor):
212
+ new_average = self.momentum * self.running_average
213
+ self.running_average = update_value + new_average
214
+
215
+ def __repr__(self) -> str:
216
+ """
217
+ Returns a string representation showing momentum, shape, statistics, and a slice of the running_average.
218
+ """
219
+ if isinstance(self.running_average, torch.Tensor):
220
+ shape = tuple(self.running_average.shape)
221
+
222
+ # Calculate statistics
223
+ with torch.no_grad():
224
+ stats = {
225
+ "mean": self.running_average.mean().item(),
226
+ "std": self.running_average.std().item(),
227
+ "min": self.running_average.min().item(),
228
+ "max": self.running_average.max().item(),
229
+ }
230
+
231
+ # Get a slice (max 3 elements per dimension)
232
+ slice_indices = tuple(slice(None, min(3, dim)) for dim in shape)
233
+ sliced_data = self.running_average[slice_indices]
234
+
235
+ # Format the slice for display (convert to float32 for numpy compatibility with bfloat16)
236
+ slice_str = str(sliced_data.detach().float().cpu().numpy())
237
+ if len(slice_str) > 200: # Truncate if too long
238
+ slice_str = slice_str[:200] + "..."
239
+
240
+ stats_str = ", ".join([f"{k}={v:.4f}" for k, v in stats.items()])
241
+
242
+ return (
243
+ f"MomentumBuffer(\n"
244
+ f" momentum={self.momentum},\n"
245
+ f" shape={shape},\n"
246
+ f" stats=[{stats_str}],\n"
247
+ f" slice={slice_str}\n"
248
+ f")"
249
+ )
250
+ else:
251
+ return f"MomentumBuffer(momentum={self.momentum}, running_average={self.running_average})"
252
+
253
+
254
+ def update_momentum_buffer(
255
+ pred_cond: torch.Tensor,
256
+ pred_uncond: torch.Tensor,
257
+ momentum_buffer: MomentumBuffer | None = None,
258
+ ):
259
+ diff = pred_cond - pred_uncond
260
+ if momentum_buffer is not None:
261
+ momentum_buffer.update(diff)
262
+
263
+
264
+ def normalized_guidance(
265
+ pred_cond: torch.Tensor,
266
+ pred_uncond: torch.Tensor,
267
+ guidance_scale: float,
268
+ momentum_buffer: MomentumBuffer | None = None,
269
+ eta: float = 1.0,
270
+ norm_threshold: float = 0.0,
271
+ use_original_formulation: bool = False,
272
+ ):
273
+ if momentum_buffer is not None:
274
+ update_momentum_buffer(pred_cond, pred_uncond, momentum_buffer)
275
+ diff = momentum_buffer.running_average
276
+ else:
277
+ diff = pred_cond - pred_uncond
278
+
279
+ dim = [-i for i in range(1, len(diff.shape))]
280
+
281
+ if norm_threshold > 0:
282
+ ones = torch.ones_like(diff)
283
+ diff_norm = diff.norm(p=2, dim=dim, keepdim=True)
284
+ scale_factor = torch.minimum(ones, norm_threshold / diff_norm)
285
+ diff = diff * scale_factor
286
+
287
+ v0, v1 = diff.double(), pred_cond.double()
288
+ v1 = torch.nn.functional.normalize(v1, dim=dim)
289
+ v0_parallel = (v0 * v1).sum(dim=dim, keepdim=True) * v1
290
+ v0_orthogonal = v0 - v0_parallel
291
+ diff_parallel, diff_orthogonal = v0_parallel.type_as(diff), v0_orthogonal.type_as(diff)
292
+ normalized_update = diff_orthogonal + eta * diff_parallel
293
+
294
+ pred = pred_cond if use_original_formulation else pred_uncond
295
+ pred = pred + guidance_scale * normalized_update
296
+
297
+ return pred
diffusers/guiders/auto_guidance.py ADDED
@@ -0,0 +1,198 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 The HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ from __future__ import annotations
16
+
17
+ import math
18
+ from typing import TYPE_CHECKING, Any
19
+
20
+ import torch
21
+
22
+ from ..configuration_utils import register_to_config
23
+ from ..hooks import HookRegistry, LayerSkipConfig
24
+ from ..hooks.layer_skip import _apply_layer_skip_hook
25
+ from .guider_utils import BaseGuidance, GuiderOutput, rescale_noise_cfg
26
+
27
+
28
+ if TYPE_CHECKING:
29
+ from ..modular_pipelines.modular_pipeline import BlockState
30
+
31
+
32
+ class AutoGuidance(BaseGuidance):
33
+ """
34
+ AutoGuidance: https://huggingface.co/papers/2406.02507
35
+
36
+ Args:
37
+ guidance_scale (`float`, defaults to `7.5`):
38
+ The scale parameter for classifier-free guidance. Higher values result in stronger conditioning on the text
39
+ prompt, while lower values allow for more freedom in generation. Higher values may lead to saturation and
40
+ deterioration of image quality.
41
+ auto_guidance_layers (`int` or `list[int]`, *optional*):
42
+ The layer indices to apply skip layer guidance to. Can be a single integer or a list of integers. If not
43
+ provided, `skip_layer_config` must be provided.
44
+ auto_guidance_config (`LayerSkipConfig` or `list[LayerSkipConfig]`, *optional*):
45
+ The configuration for the skip layer guidance. Can be a single `LayerSkipConfig` or a list of
46
+ `LayerSkipConfig`. If not provided, `skip_layer_guidance_layers` must be provided.
47
+ dropout (`float`, *optional*):
48
+ The dropout probability for autoguidance on the enabled skip layers (either with `auto_guidance_layers` or
49
+ `auto_guidance_config`). If not provided, the dropout probability will be set to 1.0.
50
+ guidance_rescale (`float`, defaults to `0.0`):
51
+ The rescale factor applied to the noise predictions. This is used to improve image quality and fix
52
+ overexposure. Based on Section 3.4 from [Common Diffusion Noise Schedules and Sample Steps are
53
+ Flawed](https://huggingface.co/papers/2305.08891).
54
+ use_original_formulation (`bool`, defaults to `False`):
55
+ Whether to use the original formulation of classifier-free guidance as proposed in the paper. By default,
56
+ we use the diffusers-native implementation that has been in the codebase for a long time. See
57
+ [~guiders.classifier_free_guidance.ClassifierFreeGuidance] for more details.
58
+ start (`float`, defaults to `0.0`):
59
+ The fraction of the total number of denoising steps after which guidance starts.
60
+ stop (`float`, defaults to `1.0`):
61
+ The fraction of the total number of denoising steps after which guidance stops.
62
+ """
63
+
64
+ _input_predictions = ["pred_cond", "pred_uncond"]
65
+
66
+ @register_to_config
67
+ def __init__(
68
+ self,
69
+ guidance_scale: float = 7.5,
70
+ auto_guidance_layers: int | list[int] | None = None,
71
+ auto_guidance_config: LayerSkipConfig | list[LayerSkipConfig] | dict[str, Any] = None,
72
+ dropout: float | None = None,
73
+ guidance_rescale: float = 0.0,
74
+ use_original_formulation: bool = False,
75
+ start: float = 0.0,
76
+ stop: float = 1.0,
77
+ enabled: bool = True,
78
+ ):
79
+ super().__init__(start, stop, enabled)
80
+
81
+ self.guidance_scale = guidance_scale
82
+ self.auto_guidance_layers = auto_guidance_layers
83
+ self.auto_guidance_config = auto_guidance_config
84
+ self.dropout = dropout
85
+ self.guidance_rescale = guidance_rescale
86
+ self.use_original_formulation = use_original_formulation
87
+
88
+ is_layer_or_config_provided = auto_guidance_layers is not None or auto_guidance_config is not None
89
+ is_layer_and_config_provided = auto_guidance_layers is not None and auto_guidance_config is not None
90
+ if not is_layer_or_config_provided:
91
+ raise ValueError(
92
+ "Either `auto_guidance_layers` or `auto_guidance_config` must be provided to enable AutoGuidance."
93
+ )
94
+ if is_layer_and_config_provided:
95
+ raise ValueError("Only one of `auto_guidance_layers` or `auto_guidance_config` can be provided.")
96
+ if auto_guidance_config is None and dropout is None:
97
+ raise ValueError("`dropout` must be provided if `auto_guidance_layers` is provided.")
98
+
99
+ if auto_guidance_layers is not None:
100
+ if isinstance(auto_guidance_layers, int):
101
+ auto_guidance_layers = [auto_guidance_layers]
102
+ if not isinstance(auto_guidance_layers, list):
103
+ raise ValueError(
104
+ f"Expected `auto_guidance_layers` to be an int or a list of ints, but got {type(auto_guidance_layers)}."
105
+ )
106
+ auto_guidance_config = [
107
+ LayerSkipConfig(layer, fqn="auto", dropout=dropout) for layer in auto_guidance_layers
108
+ ]
109
+
110
+ if isinstance(auto_guidance_config, dict):
111
+ auto_guidance_config = LayerSkipConfig.from_dict(auto_guidance_config)
112
+
113
+ if isinstance(auto_guidance_config, LayerSkipConfig):
114
+ auto_guidance_config = [auto_guidance_config]
115
+
116
+ if not isinstance(auto_guidance_config, list):
117
+ raise ValueError(
118
+ f"Expected `auto_guidance_config` to be a LayerSkipConfig or a list of LayerSkipConfig, but got {type(auto_guidance_config)}."
119
+ )
120
+ elif isinstance(next(iter(auto_guidance_config), None), dict):
121
+ auto_guidance_config = [LayerSkipConfig.from_dict(config) for config in auto_guidance_config]
122
+
123
+ self.auto_guidance_config = auto_guidance_config
124
+ self._auto_guidance_hook_names = [f"AutoGuidance_{i}" for i in range(len(self.auto_guidance_config))]
125
+
126
+ def prepare_models(self, denoiser: torch.nn.Module) -> None:
127
+ self._count_prepared += 1
128
+ if self._is_ag_enabled() and self.is_unconditional:
129
+ for name, config in zip(self._auto_guidance_hook_names, self.auto_guidance_config):
130
+ _apply_layer_skip_hook(denoiser, config, name=name)
131
+
132
+ def cleanup_models(self, denoiser: torch.nn.Module) -> None:
133
+ if self._is_ag_enabled() and self.is_unconditional:
134
+ for name in self._auto_guidance_hook_names:
135
+ registry = HookRegistry.check_if_exists_or_initialize(denoiser)
136
+ registry.remove_hook(name, recurse=True)
137
+
138
+ def prepare_inputs(self, data: dict[str, tuple[torch.Tensor, torch.Tensor]]) -> list["BlockState"]:
139
+ tuple_indices = [0] if self.num_conditions == 1 else [0, 1]
140
+ data_batches = []
141
+ for tuple_idx, input_prediction in zip(tuple_indices, self._input_predictions):
142
+ data_batch = self._prepare_batch(data, tuple_idx, input_prediction)
143
+ data_batches.append(data_batch)
144
+ return data_batches
145
+
146
+ def prepare_inputs_from_block_state(
147
+ self, data: "BlockState", input_fields: dict[str, str | tuple[str, str]]
148
+ ) -> list["BlockState"]:
149
+ tuple_indices = [0] if self.num_conditions == 1 else [0, 1]
150
+ data_batches = []
151
+ for tuple_idx, input_prediction in zip(tuple_indices, self._input_predictions):
152
+ data_batch = self._prepare_batch_from_block_state(input_fields, data, tuple_idx, input_prediction)
153
+ data_batches.append(data_batch)
154
+ return data_batches
155
+
156
+ def forward(self, pred_cond: torch.Tensor, pred_uncond: torch.Tensor | None = None) -> GuiderOutput:
157
+ pred = None
158
+
159
+ if not self._is_ag_enabled():
160
+ pred = pred_cond
161
+ else:
162
+ shift = pred_cond - pred_uncond
163
+ pred = pred_cond if self.use_original_formulation else pred_uncond
164
+ pred = pred + self.guidance_scale * shift
165
+
166
+ if self.guidance_rescale > 0.0:
167
+ pred = rescale_noise_cfg(pred, pred_cond, self.guidance_rescale)
168
+
169
+ return GuiderOutput(pred=pred, pred_cond=pred_cond, pred_uncond=pred_uncond)
170
+
171
+ @property
172
+ def is_conditional(self) -> bool:
173
+ return self._count_prepared == 1
174
+
175
+ @property
176
+ def num_conditions(self) -> int:
177
+ num_conditions = 1
178
+ if self._is_ag_enabled():
179
+ num_conditions += 1
180
+ return num_conditions
181
+
182
+ def _is_ag_enabled(self) -> bool:
183
+ if not self._enabled:
184
+ return False
185
+
186
+ is_within_range = True
187
+ if self._num_inference_steps is not None:
188
+ skip_start_step = int(self._start * self._num_inference_steps)
189
+ skip_stop_step = int(self._stop * self._num_inference_steps)
190
+ is_within_range = skip_start_step <= self._step < skip_stop_step
191
+
192
+ is_close = False
193
+ if self.use_original_formulation:
194
+ is_close = math.isclose(self.guidance_scale, 0.0)
195
+ else:
196
+ is_close = math.isclose(self.guidance_scale, 1.0)
197
+
198
+ return is_within_range and not is_close
diffusers/guiders/classifier_free_guidance.py ADDED
@@ -0,0 +1,156 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 The HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ from __future__ import annotations
16
+
17
+ import math
18
+ from typing import TYPE_CHECKING
19
+
20
+ import torch
21
+
22
+ from ..configuration_utils import register_to_config
23
+ from .guider_utils import BaseGuidance, GuiderOutput, rescale_noise_cfg
24
+
25
+
26
+ if TYPE_CHECKING:
27
+ from ..modular_pipelines.modular_pipeline import BlockState
28
+
29
+
30
+ class ClassifierFreeGuidance(BaseGuidance):
31
+ """
32
+ Implements Classifier-Free Guidance (CFG) for diffusion models.
33
+
34
+ Reference: https://huggingface.co/papers/2207.12598
35
+
36
+ CFG improves generation quality and prompt adherence by jointly training models on both conditional and
37
+ unconditional data, then combining predictions during inference. This allows trading off between quality (high
38
+ guidance) and diversity (low guidance).
39
+
40
+ **Two CFG Formulations:**
41
+
42
+ 1. **Original formulation** (from paper):
43
+ ```
44
+ x_pred = x_cond + guidance_scale * (x_cond - x_uncond)
45
+ ```
46
+ Moves conditional predictions further from unconditional ones.
47
+
48
+ 2. **Diffusers-native formulation** (default, from Imagen paper):
49
+ ```
50
+ x_pred = x_uncond + guidance_scale * (x_cond - x_uncond)
51
+ ```
52
+ Moves unconditional predictions toward conditional ones, effectively suppressing negative features (e.g., "bad
53
+ quality", "watermarks"). Equivalent in theory but more intuitive.
54
+
55
+ Use `use_original_formulation=True` to switch to the original formulation.
56
+
57
+ Args:
58
+ guidance_scale (`float`, defaults to `7.5`):
59
+ CFG scale applied by this guider during post-processing. Higher values = stronger prompt conditioning but
60
+ may reduce quality. Typical range: 1.0-20.0.
61
+ guidance_rescale (`float`, defaults to `0.0`):
62
+ Rescaling factor to prevent overexposure from high guidance scales. Based on [Common Diffusion Noise
63
+ Schedules and Sample Steps are Flawed](https://huggingface.co/papers/2305.08891). Range: 0.0 (no rescaling)
64
+ to 1.0 (full rescaling).
65
+ use_original_formulation (`bool`, defaults to `False`):
66
+ If `True`, uses the original CFG formulation from the paper. If `False` (default), uses the
67
+ diffusers-native formulation from the Imagen paper.
68
+ start (`float`, defaults to `0.0`):
69
+ Fraction of denoising steps (0.0-1.0) after which CFG starts. Use > 0.0 to disable CFG in early denoising
70
+ steps.
71
+ stop (`float`, defaults to `1.0`):
72
+ Fraction of denoising steps (0.0-1.0) after which CFG stops. Use < 1.0 to disable CFG in late denoising
73
+ steps.
74
+ enabled (`bool`, defaults to `True`):
75
+ Whether CFG is enabled. Set to `False` to disable CFG entirely (uses only conditional predictions).
76
+ """
77
+
78
+ _input_predictions = ["pred_cond", "pred_uncond"]
79
+
80
+ @register_to_config
81
+ def __init__(
82
+ self,
83
+ guidance_scale: float = 7.5,
84
+ guidance_rescale: float = 0.0,
85
+ use_original_formulation: bool = False,
86
+ start: float = 0.0,
87
+ stop: float = 1.0,
88
+ enabled: bool = True,
89
+ ):
90
+ super().__init__(start, stop, enabled)
91
+
92
+ self.guidance_scale = guidance_scale
93
+ self.guidance_rescale = guidance_rescale
94
+ self.use_original_formulation = use_original_formulation
95
+
96
+ def prepare_inputs(self, data: dict[str, tuple[torch.Tensor, torch.Tensor]]) -> list["BlockState"]:
97
+ tuple_indices = [0] if self.num_conditions == 1 else [0, 1]
98
+ data_batches = []
99
+ for tuple_idx, input_prediction in zip(tuple_indices, self._input_predictions):
100
+ data_batch = self._prepare_batch(data, tuple_idx, input_prediction)
101
+ data_batches.append(data_batch)
102
+ return data_batches
103
+
104
+ def prepare_inputs_from_block_state(
105
+ self, data: "BlockState", input_fields: dict[str, str | tuple[str, str]]
106
+ ) -> list["BlockState"]:
107
+ tuple_indices = [0] if self.num_conditions == 1 else [0, 1]
108
+ data_batches = []
109
+ for tuple_idx, input_prediction in zip(tuple_indices, self._input_predictions):
110
+ data_batch = self._prepare_batch_from_block_state(input_fields, data, tuple_idx, input_prediction)
111
+ data_batches.append(data_batch)
112
+ return data_batches
113
+
114
+ def forward(self, pred_cond: torch.Tensor, pred_uncond: torch.Tensor | None = None) -> GuiderOutput:
115
+ pred = None
116
+
117
+ if not self._is_cfg_enabled():
118
+ pred = pred_cond
119
+ else:
120
+ shift = pred_cond - pred_uncond
121
+ pred = pred_cond if self.use_original_formulation else pred_uncond
122
+ pred = pred + self.guidance_scale * shift
123
+
124
+ if self.guidance_rescale > 0.0:
125
+ pred = rescale_noise_cfg(pred, pred_cond, self.guidance_rescale)
126
+
127
+ return GuiderOutput(pred=pred, pred_cond=pred_cond, pred_uncond=pred_uncond)
128
+
129
+ @property
130
+ def is_conditional(self) -> bool:
131
+ return self._count_prepared == 1
132
+
133
+ @property
134
+ def num_conditions(self) -> int:
135
+ num_conditions = 1
136
+ if self._is_cfg_enabled():
137
+ num_conditions += 1
138
+ return num_conditions
139
+
140
+ def _is_cfg_enabled(self) -> bool:
141
+ if not self._enabled:
142
+ return False
143
+
144
+ is_within_range = True
145
+ if self._num_inference_steps is not None:
146
+ skip_start_step = int(self._start * self._num_inference_steps)
147
+ skip_stop_step = int(self._stop * self._num_inference_steps)
148
+ is_within_range = skip_start_step <= self._step < skip_stop_step
149
+
150
+ is_close = False
151
+ if self.use_original_formulation:
152
+ is_close = math.isclose(self.guidance_scale, 0.0)
153
+ else:
154
+ is_close = math.isclose(self.guidance_scale, 1.0)
155
+
156
+ return is_within_range and not is_close
diffusers/guiders/classifier_free_zero_star_guidance.py ADDED
@@ -0,0 +1,164 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 The HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ from __future__ import annotations
16
+
17
+ import math
18
+ from typing import TYPE_CHECKING
19
+
20
+ import torch
21
+
22
+ from ..configuration_utils import register_to_config
23
+ from .guider_utils import BaseGuidance, GuiderOutput, rescale_noise_cfg
24
+
25
+
26
+ if TYPE_CHECKING:
27
+ from ..modular_pipelines.modular_pipeline import BlockState
28
+
29
+
30
+ class ClassifierFreeZeroStarGuidance(BaseGuidance):
31
+ """
32
+ Classifier-free Zero* (CFG-Zero*): https://huggingface.co/papers/2503.18886
33
+
34
+ This is an implementation of the Classifier-Free Zero* guidance technique, which is a variant of classifier-free
35
+ guidance. It proposes zero initialization of the noise predictions for the first few steps of the diffusion
36
+ process, and also introduces an optimal rescaling factor for the noise predictions, which can help in improving the
37
+ quality of generated images.
38
+
39
+ The authors of the paper suggest setting zero initialization in the first 4% of the inference steps.
40
+
41
+ Args:
42
+ guidance_scale (`float`, defaults to `7.5`):
43
+ The scale parameter for classifier-free guidance. Higher values result in stronger conditioning on the text
44
+ prompt, while lower values allow for more freedom in generation. Higher values may lead to saturation and
45
+ deterioration of image quality.
46
+ zero_init_steps (`int`, defaults to `1`):
47
+ The number of inference steps for which the noise predictions are zeroed out (see Section 4.2).
48
+ guidance_rescale (`float`, defaults to `0.0`):
49
+ The rescale factor applied to the noise predictions. This is used to improve image quality and fix
50
+ overexposure. Based on Section 3.4 from [Common Diffusion Noise Schedules and Sample Steps are
51
+ Flawed](https://huggingface.co/papers/2305.08891).
52
+ use_original_formulation (`bool`, defaults to `False`):
53
+ Whether to use the original formulation of classifier-free guidance as proposed in the paper. By default,
54
+ we use the diffusers-native implementation that has been in the codebase for a long time. See
55
+ [~guiders.classifier_free_guidance.ClassifierFreeGuidance] for more details.
56
+ start (`float`, defaults to `0.01`):
57
+ The fraction of the total number of denoising steps after which guidance starts.
58
+ stop (`float`, defaults to `0.2`):
59
+ The fraction of the total number of denoising steps after which guidance stops.
60
+ """
61
+
62
+ _input_predictions = ["pred_cond", "pred_uncond"]
63
+
64
+ @register_to_config
65
+ def __init__(
66
+ self,
67
+ guidance_scale: float = 7.5,
68
+ zero_init_steps: int = 1,
69
+ guidance_rescale: float = 0.0,
70
+ use_original_formulation: bool = False,
71
+ start: float = 0.0,
72
+ stop: float = 1.0,
73
+ enabled: bool = True,
74
+ ):
75
+ super().__init__(start, stop, enabled)
76
+
77
+ self.guidance_scale = guidance_scale
78
+ self.zero_init_steps = zero_init_steps
79
+ self.guidance_rescale = guidance_rescale
80
+ self.use_original_formulation = use_original_formulation
81
+
82
+ def prepare_inputs(self, data: dict[str, tuple[torch.Tensor, torch.Tensor]]) -> list["BlockState"]:
83
+ tuple_indices = [0] if self.num_conditions == 1 else [0, 1]
84
+ data_batches = []
85
+ for tuple_idx, input_prediction in zip(tuple_indices, self._input_predictions):
86
+ data_batch = self._prepare_batch(data, tuple_idx, input_prediction)
87
+ data_batches.append(data_batch)
88
+ return data_batches
89
+
90
+ def prepare_inputs_from_block_state(
91
+ self, data: "BlockState", input_fields: dict[str, str | tuple[str, str]]
92
+ ) -> list["BlockState"]:
93
+ tuple_indices = [0] if self.num_conditions == 1 else [0, 1]
94
+ data_batches = []
95
+ for tuple_idx, input_prediction in zip(tuple_indices, self._input_predictions):
96
+ data_batch = self._prepare_batch_from_block_state(input_fields, data, tuple_idx, input_prediction)
97
+ data_batches.append(data_batch)
98
+ return data_batches
99
+
100
+ def forward(self, pred_cond: torch.Tensor, pred_uncond: torch.Tensor | None = None) -> GuiderOutput:
101
+ pred = None
102
+
103
+ # YiYi Notes: add default behavior for self._enabled == False
104
+ if not self._enabled:
105
+ pred = pred_cond
106
+
107
+ elif self._step < self.zero_init_steps:
108
+ pred = torch.zeros_like(pred_cond)
109
+ elif not self._is_cfg_enabled():
110
+ pred = pred_cond
111
+ else:
112
+ pred_cond_flat = pred_cond.flatten(1)
113
+ pred_uncond_flat = pred_uncond.flatten(1)
114
+ alpha = cfg_zero_star_scale(pred_cond_flat, pred_uncond_flat)
115
+ alpha = alpha.view(-1, *(1,) * (len(pred_cond.shape) - 1))
116
+ pred_uncond = pred_uncond * alpha
117
+ shift = pred_cond - pred_uncond
118
+ pred = pred_cond if self.use_original_formulation else pred_uncond
119
+ pred = pred + self.guidance_scale * shift
120
+
121
+ if self.guidance_rescale > 0.0:
122
+ pred = rescale_noise_cfg(pred, pred_cond, self.guidance_rescale)
123
+
124
+ return GuiderOutput(pred=pred, pred_cond=pred_cond, pred_uncond=pred_uncond)
125
+
126
+ @property
127
+ def is_conditional(self) -> bool:
128
+ return self._count_prepared == 1
129
+
130
+ @property
131
+ def num_conditions(self) -> int:
132
+ num_conditions = 1
133
+ if self._is_cfg_enabled():
134
+ num_conditions += 1
135
+ return num_conditions
136
+
137
+ def _is_cfg_enabled(self) -> bool:
138
+ if not self._enabled:
139
+ return False
140
+
141
+ is_within_range = True
142
+ if self._num_inference_steps is not None:
143
+ skip_start_step = int(self._start * self._num_inference_steps)
144
+ skip_stop_step = int(self._stop * self._num_inference_steps)
145
+ is_within_range = skip_start_step <= self._step < skip_stop_step
146
+
147
+ is_close = False
148
+ if self.use_original_formulation:
149
+ is_close = math.isclose(self.guidance_scale, 0.0)
150
+ else:
151
+ is_close = math.isclose(self.guidance_scale, 1.0)
152
+
153
+ return is_within_range and not is_close
154
+
155
+
156
+ def cfg_zero_star_scale(cond: torch.Tensor, uncond: torch.Tensor, eps: float = 1e-8) -> torch.Tensor:
157
+ cond_dtype = cond.dtype
158
+ cond = cond.float()
159
+ uncond = uncond.float()
160
+ dot_product = torch.sum(cond * uncond, dim=1, keepdim=True)
161
+ squared_norm = torch.sum(uncond**2, dim=1, keepdim=True) + eps
162
+ # st_star = v_cond^T * v_uncond / ||v_uncond||^2
163
+ scale = dot_product / squared_norm
164
+ return scale.to(dtype=cond_dtype)
diffusers/guiders/frequency_decoupled_guidance.py ADDED
@@ -0,0 +1,335 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 The HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ from __future__ import annotations
16
+
17
+ import math
18
+ from typing import TYPE_CHECKING
19
+
20
+ import torch
21
+
22
+ from ..configuration_utils import register_to_config
23
+ from ..utils import is_kornia_available
24
+ from .guider_utils import BaseGuidance, GuiderOutput, rescale_noise_cfg
25
+
26
+
27
+ if TYPE_CHECKING:
28
+ from ..modular_pipelines.modular_pipeline import BlockState
29
+
30
+
31
+ _CAN_USE_KORNIA = is_kornia_available()
32
+
33
+
34
+ if _CAN_USE_KORNIA:
35
+ from kornia.geometry import pyrup as upsample_and_blur_func
36
+ from kornia.geometry.transform import build_laplacian_pyramid as build_laplacian_pyramid_func
37
+ else:
38
+ upsample_and_blur_func = None
39
+ build_laplacian_pyramid_func = None
40
+
41
+
42
+ def project(v0: torch.Tensor, v1: torch.Tensor, upcast_to_double: bool = True) -> tuple[torch.Tensor, torch.Tensor]:
43
+ """
44
+ Project vector v0 onto vector v1, returning the parallel and orthogonal components of v0. Implementation from paper
45
+ (Algorithm 2).
46
+ """
47
+ # v0 shape: [B, ...]
48
+ # v1 shape: [B, ...]
49
+ # Assume first dim is a batch dim and all other dims are channel or "spatial" dims
50
+ all_dims_but_first = list(range(1, len(v0.shape)))
51
+ if upcast_to_double:
52
+ dtype = v0.dtype
53
+ v0, v1 = v0.double(), v1.double()
54
+ v1 = torch.nn.functional.normalize(v1, dim=all_dims_but_first)
55
+ v0_parallel = (v0 * v1).sum(dim=all_dims_but_first, keepdim=True) * v1
56
+ v0_orthogonal = v0 - v0_parallel
57
+ if upcast_to_double:
58
+ v0_parallel = v0_parallel.to(dtype)
59
+ v0_orthogonal = v0_orthogonal.to(dtype)
60
+ return v0_parallel, v0_orthogonal
61
+
62
+
63
+ def build_image_from_pyramid(pyramid: list[torch.Tensor]) -> torch.Tensor:
64
+ """
65
+ Recovers the data space latents from the Laplacian pyramid frequency space. Implementation from the paper
66
+ (Algorithm 2).
67
+ """
68
+ # pyramid shapes: [[B, C, H, W], [B, C, H/2, W/2], ...]
69
+ img = pyramid[-1]
70
+ for i in range(len(pyramid) - 2, -1, -1):
71
+ img = upsample_and_blur_func(img) + pyramid[i]
72
+ return img
73
+
74
+
75
+ class FrequencyDecoupledGuidance(BaseGuidance):
76
+ """
77
+ Frequency-Decoupled Guidance (FDG): https://huggingface.co/papers/2506.19713
78
+
79
+ FDG is a technique similar to (and based on) classifier-free guidance (CFG) which is used to improve generation
80
+ quality and condition-following in diffusion models. Like CFG, during training we jointly train the model on both
81
+ conditional and unconditional data, and use a combination of the two during inference. (If you want more details on
82
+ how CFG works, you can check out the CFG guider.)
83
+
84
+ FDG differs from CFG in that the normal CFG prediction is instead decoupled into low- and high-frequency components
85
+ using a frequency transform (such as a Laplacian pyramid). The CFG update is then performed in frequency space
86
+ separately for the low- and high-frequency components with different guidance scales. Finally, the inverse
87
+ frequency transform is used to map the CFG frequency predictions back to data space (e.g. pixel space for images)
88
+ to form the final FDG prediction.
89
+
90
+ For images, the FDG authors found that using low guidance scales for the low-frequency components retains sample
91
+ diversity and realistic color composition, while using high guidance scales for high-frequency components enhances
92
+ sample quality (such as better visual details). Therefore, they recommend using low guidance scales (low w_low) for
93
+ the low-frequency components and high guidance scales (high w_high) for the high-frequency components. As an
94
+ example, they suggest w_low = 5.0 and w_high = 10.0 for Stable Diffusion XL (see Table 8 in the paper).
95
+
96
+ As with CFG, Diffusers implements the scaling and shifting on the unconditional prediction based on the [Imagen
97
+ paper](https://huggingface.co/papers/2205.11487), which is equivalent to what the original CFG paper proposed in
98
+ theory. [x_pred = x_uncond + scale * (x_cond - x_uncond)]
99
+
100
+ The `use_original_formulation` argument can be set to `True` to use the original CFG formulation mentioned in the
101
+ paper. By default, we use the diffusers-native implementation that has been in the codebase for a long time.
102
+
103
+ Args:
104
+ guidance_scales (`list[float]`, defaults to `[10.0, 5.0]`):
105
+ The scale parameter for frequency-decoupled guidance for each frequency component, listed from highest
106
+ frequency level to lowest. Higher values result in stronger conditioning on the text prompt, while lower
107
+ values allow for more freedom in generation. Higher values may lead to saturation and deterioration of
108
+ image quality. The FDG authors recommend using higher guidance scales for higher frequency components and
109
+ lower guidance scales for lower frequency components (so `guidance_scales` should typically be sorted in
110
+ descending order).
111
+ guidance_rescale (`float` or `list[float]`, defaults to `0.0`):
112
+ The rescale factor applied to the noise predictions. This is used to improve image quality and fix
113
+ overexposure. Based on Section 3.4 from [Common Diffusion Noise Schedules and Sample Steps are
114
+ Flawed](https://huggingface.co/papers/2305.08891). If a list is supplied, it should be the same length as
115
+ `guidance_scales`.
116
+ parallel_weights (`float` or `list[float]`, *optional*):
117
+ Optional weights for the parallel component of each frequency component of the projected CFG shift. If not
118
+ set, the weights will default to `1.0` for all components, which corresponds to using the normal CFG shift
119
+ (that is, equal weights for the parallel and orthogonal components). If set, a value in `[0, 1]` is
120
+ recommended. If a list is supplied, it should be the same length as `guidance_scales`.
121
+ use_original_formulation (`bool`, defaults to `False`):
122
+ Whether to use the original formulation of classifier-free guidance as proposed in the paper. By default,
123
+ we use the diffusers-native implementation that has been in the codebase for a long time. See
124
+ [~guiders.classifier_free_guidance.ClassifierFreeGuidance] for more details.
125
+ start (`float` or `list[float]`, defaults to `0.0`):
126
+ The fraction of the total number of denoising steps after which guidance starts. If a list is supplied, it
127
+ should be the same length as `guidance_scales`.
128
+ stop (`float` or `list[float]`, defaults to `1.0`):
129
+ The fraction of the total number of denoising steps after which guidance stops. If a list is supplied, it
130
+ should be the same length as `guidance_scales`.
131
+ guidance_rescale_space (`str`, defaults to `"data"`):
132
+ Whether to performance guidance rescaling in `"data"` space (after the full FDG update in data space) or in
133
+ `"freq"` space (right after the CFG update, for each freq level). Note that frequency space rescaling is
134
+ speculative and may not produce expected results. If `"data"` is set, the first `guidance_rescale` value
135
+ will be used; otherwise, per-frequency-level guidance rescale values will be used if available.
136
+ upcast_to_double (`bool`, defaults to `True`):
137
+ Whether to upcast certain operations, such as the projection operation when using `parallel_weights`, to
138
+ float64 when performing guidance. This may result in better performance at the cost of increased runtime.
139
+ """
140
+
141
+ _input_predictions = ["pred_cond", "pred_uncond"]
142
+
143
+ @register_to_config
144
+ def __init__(
145
+ self,
146
+ guidance_scales: list[float] | tuple[float] = [10.0, 5.0],
147
+ guidance_rescale: float | list[float] | tuple[float] = 0.0,
148
+ parallel_weights: float | list[float] | tuple[float] | None = None,
149
+ use_original_formulation: bool = False,
150
+ start: float | list[float] | tuple[float] = 0.0,
151
+ stop: float | list[float] | tuple[float] = 1.0,
152
+ guidance_rescale_space: str = "data",
153
+ upcast_to_double: bool = True,
154
+ enabled: bool = True,
155
+ ):
156
+ if not _CAN_USE_KORNIA:
157
+ raise ImportError(
158
+ "The `FrequencyDecoupledGuidance` guider cannot be instantiated because the `kornia` library on which "
159
+ "it depends is not available in the current environment. You can install `kornia` with `pip install "
160
+ "kornia`."
161
+ )
162
+
163
+ # Set start to earliest start for any freq component and stop to latest stop for any freq component
164
+ min_start = start if isinstance(start, float) else min(start)
165
+ max_stop = stop if isinstance(stop, float) else max(stop)
166
+ super().__init__(min_start, max_stop, enabled)
167
+
168
+ self.guidance_scales = guidance_scales
169
+ self.levels = len(guidance_scales)
170
+
171
+ if isinstance(guidance_rescale, float):
172
+ self.guidance_rescale = [guidance_rescale] * self.levels
173
+ elif len(guidance_rescale) == self.levels:
174
+ self.guidance_rescale = guidance_rescale
175
+ else:
176
+ raise ValueError(
177
+ f"`guidance_rescale` has length {len(guidance_rescale)} but should have the same length as "
178
+ f"`guidance_scales` ({len(self.guidance_scales)})"
179
+ )
180
+ # Whether to perform guidance rescaling in frequency space (right after the CFG update) or data space (after
181
+ # transforming from frequency space back to data space)
182
+ if guidance_rescale_space not in ["data", "freq"]:
183
+ raise ValueError(
184
+ f"Guidance rescale space is {guidance_rescale_space} but must be one of `data` or `freq`."
185
+ )
186
+ self.guidance_rescale_space = guidance_rescale_space
187
+
188
+ if parallel_weights is None:
189
+ # Use normal CFG shift (equal weights for parallel and orthogonal components)
190
+ self.parallel_weights = [1.0] * self.levels
191
+ elif isinstance(parallel_weights, float):
192
+ self.parallel_weights = [parallel_weights] * self.levels
193
+ elif len(parallel_weights) == self.levels:
194
+ self.parallel_weights = parallel_weights
195
+ else:
196
+ raise ValueError(
197
+ f"`parallel_weights` has length {len(parallel_weights)} but should have the same length as "
198
+ f"`guidance_scales` ({len(self.guidance_scales)})"
199
+ )
200
+
201
+ self.use_original_formulation = use_original_formulation
202
+ self.upcast_to_double = upcast_to_double
203
+
204
+ if isinstance(start, float):
205
+ self.guidance_start = [start] * self.levels
206
+ elif len(start) == self.levels:
207
+ self.guidance_start = start
208
+ else:
209
+ raise ValueError(
210
+ f"`start` has length {len(start)} but should have the same length as `guidance_scales` "
211
+ f"({len(self.guidance_scales)})"
212
+ )
213
+ if isinstance(stop, float):
214
+ self.guidance_stop = [stop] * self.levels
215
+ elif len(stop) == self.levels:
216
+ self.guidance_stop = stop
217
+ else:
218
+ raise ValueError(
219
+ f"`stop` has length {len(stop)} but should have the same length as `guidance_scales` "
220
+ f"({len(self.guidance_scales)})"
221
+ )
222
+
223
+ def prepare_inputs(self, data: dict[str, tuple[torch.Tensor, torch.Tensor]]) -> list["BlockState"]:
224
+ tuple_indices = [0] if self.num_conditions == 1 else [0, 1]
225
+ data_batches = []
226
+ for tuple_idx, input_prediction in zip(tuple_indices, self._input_predictions):
227
+ data_batch = self._prepare_batch(data, tuple_idx, input_prediction)
228
+ data_batches.append(data_batch)
229
+ return data_batches
230
+
231
+ def prepare_inputs_from_block_state(
232
+ self, data: "BlockState", input_fields: dict[str, str | tuple[str, str]]
233
+ ) -> list["BlockState"]:
234
+ tuple_indices = [0] if self.num_conditions == 1 else [0, 1]
235
+ data_batches = []
236
+ for tuple_idx, input_prediction in zip(tuple_indices, self._input_predictions):
237
+ data_batch = self._prepare_batch_from_block_state(input_fields, data, tuple_idx, input_prediction)
238
+ data_batches.append(data_batch)
239
+ return data_batches
240
+
241
+ def forward(self, pred_cond: torch.Tensor, pred_uncond: torch.Tensor | None = None) -> GuiderOutput:
242
+ pred = None
243
+
244
+ if not self._is_fdg_enabled():
245
+ pred = pred_cond
246
+ else:
247
+ # Apply the frequency transform (e.g. Laplacian pyramid) to the conditional and unconditional predictions.
248
+ pred_cond_pyramid = build_laplacian_pyramid_func(pred_cond, self.levels)
249
+ pred_uncond_pyramid = build_laplacian_pyramid_func(pred_uncond, self.levels)
250
+
251
+ # From high frequencies to low frequencies, following the paper implementation
252
+ pred_guided_pyramid = []
253
+ parameters = zip(self.guidance_scales, self.parallel_weights, self.guidance_rescale)
254
+ for level, (guidance_scale, parallel_weight, guidance_rescale) in enumerate(parameters):
255
+ if self._is_fdg_enabled_for_level(level):
256
+ # Get the cond/uncond preds (in freq space) at the current frequency level
257
+ pred_cond_freq = pred_cond_pyramid[level]
258
+ pred_uncond_freq = pred_uncond_pyramid[level]
259
+
260
+ shift = pred_cond_freq - pred_uncond_freq
261
+
262
+ # Apply parallel weights, if used (1.0 corresponds to using the normal CFG shift)
263
+ if not math.isclose(parallel_weight, 1.0):
264
+ shift_parallel, shift_orthogonal = project(shift, pred_cond_freq, self.upcast_to_double)
265
+ shift = parallel_weight * shift_parallel + shift_orthogonal
266
+
267
+ # Apply CFG update for the current frequency level
268
+ pred = pred_cond_freq if self.use_original_formulation else pred_uncond_freq
269
+ pred = pred + guidance_scale * shift
270
+
271
+ if self.guidance_rescale_space == "freq" and guidance_rescale > 0.0:
272
+ pred = rescale_noise_cfg(pred, pred_cond_freq, guidance_rescale)
273
+
274
+ # Add the current FDG guided level to the FDG prediction pyramid
275
+ pred_guided_pyramid.append(pred)
276
+ else:
277
+ # Add the current pred_cond_pyramid level as the "non-FDG" prediction
278
+ pred_guided_pyramid.append(pred_cond_freq)
279
+
280
+ # Convert from frequency space back to data (e.g. pixel) space by applying inverse freq transform
281
+ pred = build_image_from_pyramid(pred_guided_pyramid)
282
+
283
+ # If rescaling in data space, use the first elem of self.guidance_rescale as the "global" rescale value
284
+ # across all freq levels
285
+ if self.guidance_rescale_space == "data" and self.guidance_rescale[0] > 0.0:
286
+ pred = rescale_noise_cfg(pred, pred_cond, self.guidance_rescale[0])
287
+
288
+ return GuiderOutput(pred=pred, pred_cond=pred_cond, pred_uncond=pred_uncond)
289
+
290
+ @property
291
+ def is_conditional(self) -> bool:
292
+ return self._count_prepared == 1
293
+
294
+ @property
295
+ def num_conditions(self) -> int:
296
+ num_conditions = 1
297
+ if self._is_fdg_enabled():
298
+ num_conditions += 1
299
+ return num_conditions
300
+
301
+ def _is_fdg_enabled(self) -> bool:
302
+ if not self._enabled:
303
+ return False
304
+
305
+ is_within_range = True
306
+ if self._num_inference_steps is not None:
307
+ skip_start_step = int(self._start * self._num_inference_steps)
308
+ skip_stop_step = int(self._stop * self._num_inference_steps)
309
+ is_within_range = skip_start_step <= self._step < skip_stop_step
310
+
311
+ is_close = False
312
+ if self.use_original_formulation:
313
+ is_close = all(math.isclose(guidance_scale, 0.0) for guidance_scale in self.guidance_scales)
314
+ else:
315
+ is_close = all(math.isclose(guidance_scale, 1.0) for guidance_scale in self.guidance_scales)
316
+
317
+ return is_within_range and not is_close
318
+
319
+ def _is_fdg_enabled_for_level(self, level: int) -> bool:
320
+ if not self._enabled:
321
+ return False
322
+
323
+ is_within_range = True
324
+ if self._num_inference_steps is not None:
325
+ skip_start_step = int(self.guidance_start[level] * self._num_inference_steps)
326
+ skip_stop_step = int(self.guidance_stop[level] * self._num_inference_steps)
327
+ is_within_range = skip_start_step <= self._step < skip_stop_step
328
+
329
+ is_close = False
330
+ if self.use_original_formulation:
331
+ is_close = math.isclose(self.guidance_scales[level], 0.0)
332
+ else:
333
+ is_close = math.isclose(self.guidance_scales[level], 1.0)
334
+
335
+ return is_within_range and not is_close
diffusers/guiders/guider_utils.py ADDED
@@ -0,0 +1,396 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 The HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ from __future__ import annotations
16
+
17
+ import os
18
+ from typing import TYPE_CHECKING, Any
19
+
20
+ import torch
21
+ from huggingface_hub.utils import validate_hf_hub_args
22
+ from typing_extensions import Self
23
+
24
+ from ..configuration_utils import ConfigMixin
25
+ from ..utils import BaseOutput, PushToHubMixin, get_logger
26
+
27
+
28
+ if TYPE_CHECKING:
29
+ from ..modular_pipelines.modular_pipeline import BlockState
30
+
31
+
32
+ GUIDER_CONFIG_NAME = "guider_config.json"
33
+
34
+
35
+ logger = get_logger(__name__) # pylint: disable=invalid-name
36
+
37
+
38
+ class BaseGuidance(ConfigMixin, PushToHubMixin):
39
+ r"""Base class providing the skeleton for implementing guidance techniques."""
40
+
41
+ config_name = GUIDER_CONFIG_NAME
42
+ _input_predictions = None
43
+ _identifier_key = "__guidance_identifier__"
44
+
45
+ def __init__(self, start: float = 0.0, stop: float = 1.0, enabled: bool = True):
46
+ logger.warning(
47
+ "Guiders are currently an experimental feature under active development. The API is subject to breaking changes in future releases."
48
+ )
49
+
50
+ self._start = start
51
+ self._stop = stop
52
+ self._step: int = None
53
+ self._num_inference_steps: int = None
54
+ self._timestep: torch.LongTensor = None
55
+ self._count_prepared = 0
56
+ self._input_fields: dict[str, str | tuple[str, str]] = None
57
+ self._enabled = enabled
58
+
59
+ if not (0.0 <= start < 1.0):
60
+ raise ValueError(f"Expected `start` to be between 0.0 and 1.0, but got {start}.")
61
+ if not (start <= stop <= 1.0):
62
+ raise ValueError(f"Expected `stop` to be between {start} and 1.0, but got {stop}.")
63
+
64
+ if self._input_predictions is None or not isinstance(self._input_predictions, list):
65
+ raise ValueError(
66
+ "`_input_predictions` must be a list of required prediction names for the guidance technique."
67
+ )
68
+
69
+ def new(self, **kwargs):
70
+ """
71
+ Creates a copy of this guider instance, optionally with modified configuration parameters.
72
+
73
+ Args:
74
+ **kwargs: Configuration parameters to override in the new instance. If no kwargs are provided,
75
+ returns an exact copy with the same configuration.
76
+
77
+ Returns:
78
+ A new guider instance with the same (or updated) configuration.
79
+
80
+ Example:
81
+ ```python
82
+ # Create a CFG guider
83
+ guider = ClassifierFreeGuidance(guidance_scale=3.5)
84
+
85
+ # Create an exact copy
86
+ same_guider = guider.new()
87
+
88
+ # Create a copy with different start step, keeping other config the same
89
+ new_guider = guider.new(guidance_scale=5)
90
+ ```
91
+ """
92
+ return self.__class__.from_config(self.config, **kwargs)
93
+
94
+ def disable(self):
95
+ self._enabled = False
96
+
97
+ def enable(self):
98
+ self._enabled = True
99
+
100
+ def set_state(self, step: int, num_inference_steps: int, timestep: torch.LongTensor) -> None:
101
+ self._step = step
102
+ self._num_inference_steps = num_inference_steps
103
+ self._timestep = timestep
104
+ self._count_prepared = 0
105
+
106
+ def get_state(self) -> dict[str, Any]:
107
+ """
108
+ Returns the current state of the guidance technique as a dictionary. The state variables will be included in
109
+ the __repr__ method. Returns:
110
+ `dict[str, Any]`: A dictionary containing the current state variables including:
111
+ - step: Current inference step
112
+ - num_inference_steps: Total number of inference steps
113
+ - timestep: Current timestep tensor
114
+ - count_prepared: Number of times prepare_models has been called
115
+ - enabled: Whether the guidance is enabled
116
+ - num_conditions: Number of conditions
117
+ """
118
+ state = {
119
+ "step": self._step,
120
+ "num_inference_steps": self._num_inference_steps,
121
+ "timestep": self._timestep,
122
+ "count_prepared": self._count_prepared,
123
+ "enabled": self._enabled,
124
+ "num_conditions": self.num_conditions,
125
+ }
126
+ return state
127
+
128
+ def __repr__(self) -> str:
129
+ """
130
+ Returns a string representation of the guidance object including both config and current state.
131
+ """
132
+ # Get ConfigMixin's __repr__
133
+ str_repr = super().__repr__()
134
+
135
+ # Get current state
136
+ state = self.get_state()
137
+
138
+ # Format each state variable on its own line with indentation
139
+ state_lines = []
140
+ for k, v in state.items():
141
+ # Convert value to string and handle multi-line values
142
+ v_str = str(v)
143
+ if "\n" in v_str:
144
+ # For multi-line values (like MomentumBuffer), indent subsequent lines
145
+ v_lines = v_str.split("\n")
146
+ v_str = v_lines[0] + "\n" + "\n".join([" " + line for line in v_lines[1:]])
147
+ state_lines.append(f" {k}: {v_str}")
148
+
149
+ state_str = "\n".join(state_lines)
150
+
151
+ return f"{str_repr}\nState:\n{state_str}"
152
+
153
+ def prepare_models(self, denoiser: torch.nn.Module) -> None:
154
+ """
155
+ Prepares the models for the guidance technique on a given batch of data. This method should be overridden in
156
+ subclasses to implement specific model preparation logic.
157
+ """
158
+ self._count_prepared += 1
159
+
160
+ def cleanup_models(self, denoiser: torch.nn.Module) -> None:
161
+ """
162
+ Cleans up the models for the guidance technique after a given batch of data. This method should be overridden
163
+ in subclasses to implement specific model cleanup logic. It is useful for removing any hooks or other stateful
164
+ modifications made during `prepare_models`.
165
+ """
166
+ pass
167
+
168
+ def prepare_inputs(self, data: "BlockState") -> list["BlockState"]:
169
+ raise NotImplementedError("BaseGuidance::prepare_inputs must be implemented in subclasses.")
170
+
171
+ def prepare_inputs_from_block_state(
172
+ self, data: "BlockState", input_fields: dict[str, str | tuple[str, str]]
173
+ ) -> list["BlockState"]:
174
+ raise NotImplementedError("BaseGuidance::prepare_inputs_from_block_state must be implemented in subclasses.")
175
+
176
+ def __call__(self, data: list["BlockState"]) -> Any:
177
+ if not all(hasattr(d, "noise_pred") for d in data):
178
+ raise ValueError("Expected all data to have `noise_pred` attribute.")
179
+ if len(data) != self.num_conditions:
180
+ raise ValueError(
181
+ f"Expected {self.num_conditions} data items, but got {len(data)}. Please check the input data."
182
+ )
183
+ forward_inputs = {getattr(d, self._identifier_key): d.noise_pred for d in data}
184
+ return self.forward(**forward_inputs)
185
+
186
+ def forward(self, *args, **kwargs) -> Any:
187
+ raise NotImplementedError("BaseGuidance::forward must be implemented in subclasses.")
188
+
189
+ @property
190
+ def is_conditional(self) -> bool:
191
+ raise NotImplementedError("BaseGuidance::is_conditional must be implemented in subclasses.")
192
+
193
+ @property
194
+ def is_unconditional(self) -> bool:
195
+ return not self.is_conditional
196
+
197
+ @property
198
+ def num_conditions(self) -> int:
199
+ raise NotImplementedError("BaseGuidance::num_conditions must be implemented in subclasses.")
200
+
201
+ @classmethod
202
+ def _prepare_batch(
203
+ cls,
204
+ data: dict[str, tuple[torch.Tensor, torch.Tensor]],
205
+ tuple_index: int,
206
+ identifier: str,
207
+ ) -> "BlockState":
208
+ """
209
+ Prepares a batch of data for the guidance technique. This method is used in the `prepare_inputs` method of the
210
+ `BaseGuidance` class. It prepares the batch based on the provided tuple index.
211
+
212
+ Args:
213
+ input_fields (`dict[str, str | tuple[str, str]]`):
214
+ A dictionary where the keys are the names of the fields that will be used to store the data once it is
215
+ prepared with `prepare_inputs`. The values can be either a string or a tuple of length 2, which is used
216
+ to look up the required data provided for preparation. If a string is provided, it will be used as the
217
+ conditional data (or unconditional if used with a guidance method that requires it). If a tuple of
218
+ length 2 is provided, the first element must be the conditional data identifier and the second element
219
+ must be the unconditional data identifier or None.
220
+ data (`BlockState`):
221
+ The input data to be prepared.
222
+ tuple_index (`int`):
223
+ The index to use when accessing input fields that are tuples.
224
+
225
+ Returns:
226
+ `BlockState`: The prepared batch of data.
227
+ """
228
+ from ..modular_pipelines.modular_pipeline import BlockState
229
+
230
+ data_batch = {}
231
+ for key, value in data.items():
232
+ try:
233
+ if isinstance(value, torch.Tensor):
234
+ data_batch[key] = value
235
+ elif isinstance(value, tuple):
236
+ data_batch[key] = value[tuple_index]
237
+ else:
238
+ raise ValueError(f"Invalid value type: {type(value)}")
239
+ except ValueError:
240
+ logger.debug(f"`data` does not have attribute(s) {value}, skipping.")
241
+ data_batch[cls._identifier_key] = identifier
242
+ return BlockState(**data_batch)
243
+
244
+ @classmethod
245
+ def _prepare_batch_from_block_state(
246
+ cls,
247
+ input_fields: dict[str, str | tuple[str, str]],
248
+ data: "BlockState",
249
+ tuple_index: int,
250
+ identifier: str,
251
+ ) -> "BlockState":
252
+ """
253
+ Prepares a batch of data for the guidance technique. This method is used in the `prepare_inputs` method of the
254
+ `BaseGuidance` class. It prepares the batch based on the provided tuple index.
255
+
256
+ Args:
257
+ input_fields (`dict[str, str | tuple[str, str]]`):
258
+ A dictionary where the keys are the names of the fields that will be used to store the data once it is
259
+ prepared with `prepare_inputs`. The values can be either a string or a tuple of length 2, which is used
260
+ to look up the required data provided for preparation. If a string is provided, it will be used as the
261
+ conditional data (or unconditional if used with a guidance method that requires it). If a tuple of
262
+ length 2 is provided, the first element must be the conditional data identifier and the second element
263
+ must be the unconditional data identifier or None.
264
+ data (`BlockState`):
265
+ The input data to be prepared.
266
+ tuple_index (`int`):
267
+ The index to use when accessing input fields that are tuples.
268
+
269
+ Returns:
270
+ `BlockState`: The prepared batch of data.
271
+ """
272
+ from ..modular_pipelines.modular_pipeline import BlockState
273
+
274
+ data_batch = {}
275
+ for key, value in input_fields.items():
276
+ try:
277
+ if isinstance(value, str):
278
+ data_batch[key] = getattr(data, value)
279
+ elif isinstance(value, tuple):
280
+ data_batch[key] = getattr(data, value[tuple_index])
281
+ else:
282
+ # We've already checked that value is a string or a tuple of strings with length 2
283
+ pass
284
+ except AttributeError:
285
+ logger.debug(f"`data` does not have attribute(s) {value}, skipping.")
286
+ data_batch[cls._identifier_key] = identifier
287
+ return BlockState(**data_batch)
288
+
289
+ @classmethod
290
+ @validate_hf_hub_args
291
+ def from_pretrained(
292
+ cls,
293
+ pretrained_model_name_or_path: str | os.PathLike | None = None,
294
+ subfolder: str | None = None,
295
+ return_unused_kwargs=False,
296
+ **kwargs,
297
+ ) -> Self:
298
+ r"""
299
+ Instantiate a guider from a pre-defined JSON configuration file in a local directory or Hub repository.
300
+
301
+ Parameters:
302
+ pretrained_model_name_or_path (`str` or `os.PathLike`, *optional*):
303
+ Can be either:
304
+
305
+ - A string, the *model id* (for example `google/ddpm-celebahq-256`) of a pretrained model hosted on
306
+ the Hub.
307
+ - A path to a *directory* (for example `./my_model_directory`) containing the guider configuration
308
+ saved with [`~BaseGuidance.save_pretrained`].
309
+ subfolder (`str`, *optional*):
310
+ The subfolder location of a model file within a larger model repository on the Hub or locally.
311
+ return_unused_kwargs (`bool`, *optional*, defaults to `False`):
312
+ Whether kwargs that are not consumed by the Python class should be returned or not.
313
+ cache_dir (`str | os.PathLike`, *optional*):
314
+ Path to a directory where a downloaded pretrained model configuration is cached if the standard cache
315
+ is not used.
316
+ force_download (`bool`, *optional*, defaults to `False`):
317
+ Whether or not to force the (re-)download of the model weights and configuration files, overriding the
318
+ cached versions if they exist.
319
+
320
+ proxies (`dict[str, str]`, *optional*):
321
+ A dictionary of proxy servers to use by protocol or endpoint, for example, `{'http': 'foo.bar:3128',
322
+ 'http://hostname': 'foo.bar:4012'}`. The proxies are used on each request.
323
+ output_loading_info(`bool`, *optional*, defaults to `False`):
324
+ Whether or not to also return a dictionary containing missing keys, unexpected keys and error messages.
325
+ local_files_only(`bool`, *optional*, defaults to `False`):
326
+ Whether to only load local model weights and configuration files or not. If set to `True`, the model
327
+ won't be downloaded from the Hub.
328
+ token (`str` or *bool*, *optional*):
329
+ The token to use as HTTP bearer authorization for remote files. If `True`, the token generated from
330
+ `diffusers-cli login` (stored in `~/.huggingface`) is used.
331
+ revision (`str`, *optional*, defaults to `"main"`):
332
+ The specific model version to use. It can be a branch name, a tag name, a commit id, or any identifier
333
+ allowed by Git.
334
+
335
+ > [!TIP] > To use private or [gated models](https://huggingface.co/docs/hub/models-gated#gated-models), log-in
336
+ with `hf > auth login`. You can also activate the special >
337
+ ["offline-mode"](https://huggingface.co/diffusers/installation.html#offline-mode) to use this method in a >
338
+ firewalled environment.
339
+
340
+ """
341
+ config, kwargs, commit_hash = cls.load_config(
342
+ pretrained_model_name_or_path=pretrained_model_name_or_path,
343
+ subfolder=subfolder,
344
+ return_unused_kwargs=True,
345
+ return_commit_hash=True,
346
+ **kwargs,
347
+ )
348
+ return cls.from_config(config, return_unused_kwargs=return_unused_kwargs, **kwargs)
349
+
350
+ def save_pretrained(self, save_directory: str | os.PathLike, push_to_hub: bool = False, **kwargs):
351
+ """
352
+ Save a guider configuration object to a directory so that it can be reloaded using the
353
+ [`~BaseGuidance.from_pretrained`] class method.
354
+
355
+ Args:
356
+ save_directory (`str` or `os.PathLike`):
357
+ Directory where the configuration JSON file will be saved (will be created if it does not exist).
358
+ push_to_hub (`bool`, *optional*, defaults to `False`):
359
+ Whether or not to push your model to the Hugging Face Hub after saving it. You can specify the
360
+ repository you want to push to with `repo_id` (will default to the name of `save_directory` in your
361
+ namespace).
362
+ kwargs (`dict[str, Any]`, *optional*):
363
+ Additional keyword arguments passed along to the [`~utils.PushToHubMixin.push_to_hub`] method.
364
+ """
365
+ self.save_config(save_directory=save_directory, push_to_hub=push_to_hub, **kwargs)
366
+
367
+
368
+ class GuiderOutput(BaseOutput):
369
+ pred: torch.Tensor
370
+ pred_cond: torch.Tensor | None
371
+ pred_uncond: torch.Tensor | None
372
+
373
+
374
+ def rescale_noise_cfg(noise_cfg, noise_pred_text, guidance_rescale=0.0):
375
+ r"""
376
+ Rescales `noise_cfg` tensor based on `guidance_rescale` to improve image quality and fix overexposure. Based on
377
+ Section 3.4 from [Common Diffusion Noise Schedules and Sample Steps are
378
+ Flawed](https://huggingface.co/papers/2305.08891).
379
+
380
+ Args:
381
+ noise_cfg (`torch.Tensor`):
382
+ The predicted noise tensor for the guided diffusion process.
383
+ noise_pred_text (`torch.Tensor`):
384
+ The predicted noise tensor for the text-guided diffusion process.
385
+ guidance_rescale (`float`, *optional*, defaults to 0.0):
386
+ A rescale factor applied to the noise predictions.
387
+ Returns:
388
+ noise_cfg (`torch.Tensor`): The rescaled noise prediction tensor.
389
+ """
390
+ std_text = noise_pred_text.std(dim=list(range(1, noise_pred_text.ndim)), keepdim=True)
391
+ std_cfg = noise_cfg.std(dim=list(range(1, noise_cfg.ndim)), keepdim=True)
392
+ # rescale the results from guidance (fixes overexposure)
393
+ noise_pred_rescaled = noise_cfg * (std_text / std_cfg)
394
+ # mix with the original results from guidance by factor guidance_rescale to avoid "plain looking" images
395
+ noise_cfg = guidance_rescale * noise_pred_rescaled + (1 - guidance_rescale) * noise_cfg
396
+ return noise_cfg
diffusers/guiders/magnitude_aware_guidance.py ADDED
@@ -0,0 +1,159 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 The HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ import math
16
+ from typing import TYPE_CHECKING
17
+
18
+ import torch
19
+
20
+ from ..configuration_utils import register_to_config
21
+ from .guider_utils import BaseGuidance, GuiderOutput, rescale_noise_cfg
22
+
23
+
24
+ if TYPE_CHECKING:
25
+ from ..modular_pipelines.modular_pipeline import BlockState
26
+
27
+
28
+ class MagnitudeAwareGuidance(BaseGuidance):
29
+ """
30
+ Magnitude-Aware Mitigation for Boosted Guidance (MAMBO-G): https://huggingface.co/papers/2508.03442
31
+
32
+ Args:
33
+ guidance_scale (`float`, defaults to `10.0`):
34
+ The scale parameter for classifier-free guidance. Higher values result in stronger conditioning on the text
35
+ prompt, while lower values allow for more freedom in generation. Higher values may lead to saturation and
36
+ deterioration of image quality.
37
+ alpha (`float`, defaults to `8.0`):
38
+ The alpha parameter for the magnitude-aware guidance. Higher values cause more aggressive supression of
39
+ guidance scale when the magnitude of the guidance update is large.
40
+ guidance_rescale (`float`, defaults to `0.0`):
41
+ The rescale factor applied to the noise predictions. This is used to improve image quality and fix
42
+ overexposure. Based on Section 3.4 from [Common Diffusion Noise Schedules and Sample Steps are
43
+ Flawed](https://huggingface.co/papers/2305.08891).
44
+ use_original_formulation (`bool`, defaults to `False`):
45
+ Whether to use the original formulation of classifier-free guidance as proposed in the paper. By default,
46
+ we use the diffusers-native implementation that has been in the codebase for a long time. See
47
+ [~guiders.classifier_free_guidance.ClassifierFreeGuidance] for more details.
48
+ start (`float`, defaults to `0.0`):
49
+ The fraction of the total number of denoising steps after which guidance starts.
50
+ stop (`float`, defaults to `1.0`):
51
+ The fraction of the total number of denoising steps after which guidance stops.
52
+ """
53
+
54
+ _input_predictions = ["pred_cond", "pred_uncond"]
55
+
56
+ @register_to_config
57
+ def __init__(
58
+ self,
59
+ guidance_scale: float = 10.0,
60
+ alpha: float = 8.0,
61
+ guidance_rescale: float = 0.0,
62
+ use_original_formulation: bool = False,
63
+ start: float = 0.0,
64
+ stop: float = 1.0,
65
+ enabled: bool = True,
66
+ ):
67
+ super().__init__(start, stop, enabled)
68
+
69
+ self.guidance_scale = guidance_scale
70
+ self.alpha = alpha
71
+ self.guidance_rescale = guidance_rescale
72
+ self.use_original_formulation = use_original_formulation
73
+
74
+ def prepare_inputs(self, data: dict[str, tuple[torch.Tensor, torch.Tensor]]) -> list["BlockState"]:
75
+ tuple_indices = [0] if self.num_conditions == 1 else [0, 1]
76
+ data_batches = []
77
+ for tuple_idx, input_prediction in zip(tuple_indices, self._input_predictions):
78
+ data_batch = self._prepare_batch(data, tuple_idx, input_prediction)
79
+ data_batches.append(data_batch)
80
+ return data_batches
81
+
82
+ def prepare_inputs_from_block_state(
83
+ self, data: "BlockState", input_fields: dict[str, str | tuple[str, str]]
84
+ ) -> list["BlockState"]:
85
+ tuple_indices = [0] if self.num_conditions == 1 else [0, 1]
86
+ data_batches = []
87
+ for tuple_idx, input_prediction in zip(tuple_indices, self._input_predictions):
88
+ data_batch = self._prepare_batch_from_block_state(input_fields, data, tuple_idx, input_prediction)
89
+ data_batches.append(data_batch)
90
+ return data_batches
91
+
92
+ def forward(self, pred_cond: torch.Tensor, pred_uncond: torch.Tensor | None = None) -> GuiderOutput:
93
+ pred = None
94
+
95
+ if not self._is_mambo_g_enabled():
96
+ pred = pred_cond
97
+ else:
98
+ pred = mambo_guidance(
99
+ pred_cond,
100
+ pred_uncond,
101
+ self.guidance_scale,
102
+ self.alpha,
103
+ self.use_original_formulation,
104
+ )
105
+
106
+ if self.guidance_rescale > 0.0:
107
+ pred = rescale_noise_cfg(pred, pred_cond, self.guidance_rescale)
108
+
109
+ return GuiderOutput(pred=pred, pred_cond=pred_cond, pred_uncond=pred_uncond)
110
+
111
+ @property
112
+ def is_conditional(self) -> bool:
113
+ return self._count_prepared == 1
114
+
115
+ @property
116
+ def num_conditions(self) -> int:
117
+ num_conditions = 1
118
+ if self._is_mambo_g_enabled():
119
+ num_conditions += 1
120
+ return num_conditions
121
+
122
+ def _is_mambo_g_enabled(self) -> bool:
123
+ if not self._enabled:
124
+ return False
125
+
126
+ is_within_range = True
127
+ if self._num_inference_steps is not None:
128
+ skip_start_step = int(self._start * self._num_inference_steps)
129
+ skip_stop_step = int(self._stop * self._num_inference_steps)
130
+ is_within_range = skip_start_step <= self._step < skip_stop_step
131
+
132
+ is_close = False
133
+ if self.use_original_formulation:
134
+ is_close = math.isclose(self.guidance_scale, 0.0)
135
+ else:
136
+ is_close = math.isclose(self.guidance_scale, 1.0)
137
+
138
+ return is_within_range and not is_close
139
+
140
+
141
+ def mambo_guidance(
142
+ pred_cond: torch.Tensor,
143
+ pred_uncond: torch.Tensor,
144
+ guidance_scale: float,
145
+ alpha: float = 8.0,
146
+ use_original_formulation: bool = False,
147
+ ):
148
+ dim = list(range(1, len(pred_cond.shape)))
149
+ diff = pred_cond - pred_uncond
150
+ ratio = torch.norm(diff, dim=dim, keepdim=True) / torch.norm(pred_uncond, dim=dim, keepdim=True)
151
+ guidance_scale_final = (
152
+ guidance_scale * torch.exp(-alpha * ratio)
153
+ if use_original_formulation
154
+ else 1.0 + (guidance_scale - 1.0) * torch.exp(-alpha * ratio)
155
+ )
156
+ pred = pred_cond if use_original_formulation else pred_uncond
157
+ pred = pred + guidance_scale_final * diff
158
+
159
+ return pred
diffusers/guiders/perturbed_attention_guidance.py ADDED
@@ -0,0 +1,289 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 The HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ from __future__ import annotations
16
+
17
+ import math
18
+ from typing import TYPE_CHECKING, Any
19
+
20
+ import torch
21
+
22
+ from ..configuration_utils import register_to_config
23
+ from ..hooks import HookRegistry, LayerSkipConfig
24
+ from ..hooks.layer_skip import _apply_layer_skip_hook
25
+ from ..utils import get_logger
26
+ from .guider_utils import BaseGuidance, GuiderOutput, rescale_noise_cfg
27
+
28
+
29
+ if TYPE_CHECKING:
30
+ from ..modular_pipelines.modular_pipeline import BlockState
31
+
32
+
33
+ logger = get_logger(__name__) # pylint: disable=invalid-name
34
+
35
+
36
+ class PerturbedAttentionGuidance(BaseGuidance):
37
+ """
38
+ Perturbed Attention Guidance (PAG): https://huggingface.co/papers/2403.17377
39
+
40
+ The intution behind PAG can be thought of as moving the CFG predicted distribution estimates further away from
41
+ worse versions of the conditional distribution estimates. PAG was one of the first techniques to introduce the idea
42
+ of using a worse version of the trained model for better guiding itself in the denoising process. It perturbs the
43
+ attention scores of the latent stream by replacing the score matrix with an identity matrix for selectively chosen
44
+ layers.
45
+
46
+ Additional reading:
47
+ - [Guiding a Diffusion Model with a Bad Version of Itself](https://huggingface.co/papers/2406.02507)
48
+
49
+ PAG is implemented with similar implementation to SkipLayerGuidance due to overlap in the configuration parameters
50
+ and implementation details.
51
+
52
+ Args:
53
+ guidance_scale (`float`, defaults to `7.5`):
54
+ The scale parameter for classifier-free guidance. Higher values result in stronger conditioning on the text
55
+ prompt, while lower values allow for more freedom in generation. Higher values may lead to saturation and
56
+ deterioration of image quality.
57
+ perturbed_guidance_scale (`float`, defaults to `2.8`):
58
+ The scale parameter for perturbed attention guidance.
59
+ perturbed_guidance_start (`float`, defaults to `0.01`):
60
+ The fraction of the total number of denoising steps after which perturbed attention guidance starts.
61
+ perturbed_guidance_stop (`float`, defaults to `0.2`):
62
+ The fraction of the total number of denoising steps after which perturbed attention guidance stops.
63
+ perturbed_guidance_layers (`int` or `list[int]`, *optional*):
64
+ The layer indices to apply perturbed attention guidance to. Can be a single integer or a list of integers.
65
+ If not provided, `perturbed_guidance_config` must be provided.
66
+ perturbed_guidance_config (`LayerSkipConfig` or `list[LayerSkipConfig]`, *optional*):
67
+ The configuration for the perturbed attention guidance. Can be a single `LayerSkipConfig` or a list of
68
+ `LayerSkipConfig`. If not provided, `perturbed_guidance_layers` must be provided.
69
+ guidance_rescale (`float`, defaults to `0.0`):
70
+ The rescale factor applied to the noise predictions. This is used to improve image quality and fix
71
+ overexposure. Based on Section 3.4 from [Common Diffusion Noise Schedules and Sample Steps are
72
+ Flawed](https://huggingface.co/papers/2305.08891).
73
+ use_original_formulation (`bool`, defaults to `False`):
74
+ Whether to use the original formulation of classifier-free guidance as proposed in the paper. By default,
75
+ we use the diffusers-native implementation that has been in the codebase for a long time. See
76
+ [~guiders.classifier_free_guidance.ClassifierFreeGuidance] for more details.
77
+ start (`float`, defaults to `0.01`):
78
+ The fraction of the total number of denoising steps after which guidance starts.
79
+ stop (`float`, defaults to `0.2`):
80
+ The fraction of the total number of denoising steps after which guidance stops.
81
+ """
82
+
83
+ # NOTE: The current implementation does not account for joint latent conditioning (text + image/video tokens in
84
+ # the same latent stream). It assumes the entire latent is a single stream of visual tokens. It would be very
85
+ # complex to support joint latent conditioning in a model-agnostic manner without specializing the implementation
86
+ # for each model architecture.
87
+
88
+ _input_predictions = ["pred_cond", "pred_uncond", "pred_cond_skip"]
89
+
90
+ @register_to_config
91
+ def __init__(
92
+ self,
93
+ guidance_scale: float = 7.5,
94
+ perturbed_guidance_scale: float = 2.8,
95
+ perturbed_guidance_start: float = 0.01,
96
+ perturbed_guidance_stop: float = 0.2,
97
+ perturbed_guidance_layers: int | list[int] | None = None,
98
+ perturbed_guidance_config: LayerSkipConfig | list[LayerSkipConfig] | dict[str, Any] = None,
99
+ guidance_rescale: float = 0.0,
100
+ use_original_formulation: bool = False,
101
+ start: float = 0.0,
102
+ stop: float = 1.0,
103
+ enabled: bool = True,
104
+ ):
105
+ super().__init__(start, stop, enabled)
106
+
107
+ self.guidance_scale = guidance_scale
108
+ self.skip_layer_guidance_scale = perturbed_guidance_scale
109
+ self.skip_layer_guidance_start = perturbed_guidance_start
110
+ self.skip_layer_guidance_stop = perturbed_guidance_stop
111
+ self.guidance_rescale = guidance_rescale
112
+ self.use_original_formulation = use_original_formulation
113
+
114
+ if perturbed_guidance_config is None:
115
+ if perturbed_guidance_layers is None:
116
+ raise ValueError(
117
+ "`perturbed_guidance_layers` must be provided if `perturbed_guidance_config` is not specified."
118
+ )
119
+ perturbed_guidance_config = LayerSkipConfig(
120
+ indices=perturbed_guidance_layers,
121
+ fqn="auto",
122
+ skip_attention=False,
123
+ skip_attention_scores=True,
124
+ skip_ff=False,
125
+ )
126
+ else:
127
+ if perturbed_guidance_layers is not None:
128
+ raise ValueError(
129
+ "`perturbed_guidance_layers` should not be provided if `perturbed_guidance_config` is specified."
130
+ )
131
+
132
+ if isinstance(perturbed_guidance_config, dict):
133
+ perturbed_guidance_config = LayerSkipConfig.from_dict(perturbed_guidance_config)
134
+
135
+ if isinstance(perturbed_guidance_config, LayerSkipConfig):
136
+ perturbed_guidance_config = [perturbed_guidance_config]
137
+
138
+ if not isinstance(perturbed_guidance_config, list):
139
+ raise ValueError(
140
+ "`perturbed_guidance_config` must be a `LayerSkipConfig`, a list of `LayerSkipConfig`, or a dict that can be converted to a `LayerSkipConfig`."
141
+ )
142
+ elif isinstance(next(iter(perturbed_guidance_config), None), dict):
143
+ perturbed_guidance_config = [LayerSkipConfig.from_dict(config) for config in perturbed_guidance_config]
144
+
145
+ for config in perturbed_guidance_config:
146
+ if config.skip_attention or not config.skip_attention_scores or config.skip_ff:
147
+ logger.warning(
148
+ "Perturbed Attention Guidance is designed to perturb attention scores, so `skip_attention` should be False, `skip_attention_scores` should be True, and `skip_ff` should be False. "
149
+ "Please check your configuration. Modifying the config to match the expected values."
150
+ )
151
+ config.skip_attention = False
152
+ config.skip_attention_scores = True
153
+ config.skip_ff = False
154
+
155
+ self.skip_layer_config = perturbed_guidance_config
156
+ self._skip_layer_hook_names = [f"SkipLayerGuidance_{i}" for i in range(len(self.skip_layer_config))]
157
+
158
+ # Copied from diffusers.guiders.skip_layer_guidance.SkipLayerGuidance.prepare_models
159
+ def prepare_models(self, denoiser: torch.nn.Module) -> None:
160
+ self._count_prepared += 1
161
+ if self._is_slg_enabled() and self.is_conditional and self._count_prepared > 1:
162
+ for name, config in zip(self._skip_layer_hook_names, self.skip_layer_config):
163
+ _apply_layer_skip_hook(denoiser, config, name=name)
164
+
165
+ # Copied from diffusers.guiders.skip_layer_guidance.SkipLayerGuidance.cleanup_models
166
+ def cleanup_models(self, denoiser: torch.nn.Module) -> None:
167
+ if self._is_slg_enabled() and self.is_conditional and self._count_prepared > 1:
168
+ registry = HookRegistry.check_if_exists_or_initialize(denoiser)
169
+ # Remove the hooks after inference
170
+ for hook_name in self._skip_layer_hook_names:
171
+ registry.remove_hook(hook_name, recurse=True)
172
+
173
+ # Copied from diffusers.guiders.skip_layer_guidance.SkipLayerGuidance.prepare_inputs
174
+ def prepare_inputs(self, data: dict[str, tuple[torch.Tensor, torch.Tensor]]) -> list["BlockState"]:
175
+ if self.num_conditions == 1:
176
+ tuple_indices = [0]
177
+ input_predictions = ["pred_cond"]
178
+ elif self.num_conditions == 2:
179
+ tuple_indices = [0, 1]
180
+ input_predictions = (
181
+ ["pred_cond", "pred_uncond"] if self._is_cfg_enabled() else ["pred_cond", "pred_cond_skip"]
182
+ )
183
+ else:
184
+ tuple_indices = [0, 1, 0]
185
+ input_predictions = ["pred_cond", "pred_uncond", "pred_cond_skip"]
186
+ data_batches = []
187
+ for tuple_idx, input_prediction in zip(tuple_indices, input_predictions):
188
+ data_batch = self._prepare_batch(data, tuple_idx, input_prediction)
189
+ data_batches.append(data_batch)
190
+ return data_batches
191
+
192
+ def prepare_inputs_from_block_state(
193
+ self, data: "BlockState", input_fields: dict[str, str | tuple[str, str]]
194
+ ) -> list["BlockState"]:
195
+ if self.num_conditions == 1:
196
+ tuple_indices = [0]
197
+ input_predictions = ["pred_cond"]
198
+ elif self.num_conditions == 2:
199
+ tuple_indices = [0, 1]
200
+ input_predictions = (
201
+ ["pred_cond", "pred_uncond"] if self._is_cfg_enabled() else ["pred_cond", "pred_cond_skip"]
202
+ )
203
+ else:
204
+ tuple_indices = [0, 1, 0]
205
+ input_predictions = ["pred_cond", "pred_uncond", "pred_cond_skip"]
206
+ data_batches = []
207
+ for tuple_idx, input_prediction in zip(tuple_indices, input_predictions):
208
+ data_batch = self._prepare_batch_from_block_state(input_fields, data, tuple_idx, input_prediction)
209
+ data_batches.append(data_batch)
210
+ return data_batches
211
+
212
+ # Copied from diffusers.guiders.skip_layer_guidance.SkipLayerGuidance.forward
213
+ def forward(
214
+ self,
215
+ pred_cond: torch.Tensor,
216
+ pred_uncond: torch.Tensor | None = None,
217
+ pred_cond_skip: torch.Tensor | None = None,
218
+ ) -> GuiderOutput:
219
+ pred = None
220
+
221
+ if not self._is_cfg_enabled() and not self._is_slg_enabled():
222
+ pred = pred_cond
223
+ elif not self._is_cfg_enabled():
224
+ shift = pred_cond - pred_cond_skip
225
+ pred = pred_cond if self.use_original_formulation else pred_cond_skip
226
+ pred = pred + self.skip_layer_guidance_scale * shift
227
+ elif not self._is_slg_enabled():
228
+ shift = pred_cond - pred_uncond
229
+ pred = pred_cond if self.use_original_formulation else pred_uncond
230
+ pred = pred + self.guidance_scale * shift
231
+ else:
232
+ shift = pred_cond - pred_uncond
233
+ shift_skip = pred_cond - pred_cond_skip
234
+ pred = pred_cond if self.use_original_formulation else pred_uncond
235
+ pred = pred + self.guidance_scale * shift + self.skip_layer_guidance_scale * shift_skip
236
+
237
+ if self.guidance_rescale > 0.0:
238
+ pred = rescale_noise_cfg(pred, pred_cond, self.guidance_rescale)
239
+
240
+ return GuiderOutput(pred=pred, pred_cond=pred_cond, pred_uncond=pred_uncond)
241
+
242
+ @property
243
+ # Copied from diffusers.guiders.skip_layer_guidance.SkipLayerGuidance.is_conditional
244
+ def is_conditional(self) -> bool:
245
+ return self._count_prepared == 1 or self._count_prepared == 3
246
+
247
+ @property
248
+ # Copied from diffusers.guiders.skip_layer_guidance.SkipLayerGuidance.num_conditions
249
+ def num_conditions(self) -> int:
250
+ num_conditions = 1
251
+ if self._is_cfg_enabled():
252
+ num_conditions += 1
253
+ if self._is_slg_enabled():
254
+ num_conditions += 1
255
+ return num_conditions
256
+
257
+ # Copied from diffusers.guiders.skip_layer_guidance.SkipLayerGuidance._is_cfg_enabled
258
+ def _is_cfg_enabled(self) -> bool:
259
+ if not self._enabled:
260
+ return False
261
+
262
+ is_within_range = True
263
+ if self._num_inference_steps is not None:
264
+ skip_start_step = int(self._start * self._num_inference_steps)
265
+ skip_stop_step = int(self._stop * self._num_inference_steps)
266
+ is_within_range = skip_start_step <= self._step < skip_stop_step
267
+
268
+ is_close = False
269
+ if self.use_original_formulation:
270
+ is_close = math.isclose(self.guidance_scale, 0.0)
271
+ else:
272
+ is_close = math.isclose(self.guidance_scale, 1.0)
273
+
274
+ return is_within_range and not is_close
275
+
276
+ # Copied from diffusers.guiders.skip_layer_guidance.SkipLayerGuidance._is_slg_enabled
277
+ def _is_slg_enabled(self) -> bool:
278
+ if not self._enabled:
279
+ return False
280
+
281
+ is_within_range = True
282
+ if self._num_inference_steps is not None:
283
+ skip_start_step = int(self.skip_layer_guidance_start * self._num_inference_steps)
284
+ skip_stop_step = int(self.skip_layer_guidance_stop * self._num_inference_steps)
285
+ is_within_range = skip_start_step < self._step < skip_stop_step
286
+
287
+ is_zero = math.isclose(self.skip_layer_guidance_scale, 0.0)
288
+
289
+ return is_within_range and not is_zero
diffusers/guiders/skip_layer_guidance.py ADDED
@@ -0,0 +1,280 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 The HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ from __future__ import annotations
16
+
17
+ import math
18
+ from typing import TYPE_CHECKING, Any
19
+
20
+ import torch
21
+
22
+ from ..configuration_utils import register_to_config
23
+ from ..hooks import HookRegistry, LayerSkipConfig
24
+ from ..hooks.layer_skip import _apply_layer_skip_hook
25
+ from .guider_utils import BaseGuidance, GuiderOutput, rescale_noise_cfg
26
+
27
+
28
+ if TYPE_CHECKING:
29
+ from ..modular_pipelines.modular_pipeline import BlockState
30
+
31
+
32
+ class SkipLayerGuidance(BaseGuidance):
33
+ """
34
+ Skip Layer Guidance (SLG): https://github.com/Stability-AI/sd3.5
35
+
36
+ Spatio-Temporal Guidance (STG): https://huggingface.co/papers/2411.18664
37
+
38
+ SLG was introduced by StabilityAI for improving structure and anotomy coherence in generated images. It works by
39
+ skipping the forward pass of specified transformer blocks during the denoising process on an additional conditional
40
+ batch of data, apart from the conditional and unconditional batches already used in CFG
41
+ ([~guiders.classifier_free_guidance.ClassifierFreeGuidance]), and then scaling and shifting the CFG predictions
42
+ based on the difference between conditional without skipping and conditional with skipping predictions.
43
+
44
+ The intution behind SLG can be thought of as moving the CFG predicted distribution estimates further away from
45
+ worse versions of the conditional distribution estimates (because skipping layers is equivalent to using a worse
46
+ version of the model for the conditional prediction).
47
+
48
+ STG is an improvement and follow-up work combining ideas from SLG, PAG and similar techniques for improving
49
+ generation quality in video diffusion models.
50
+
51
+ Additional reading:
52
+ - [Guiding a Diffusion Model with a Bad Version of Itself](https://huggingface.co/papers/2406.02507)
53
+
54
+ The values for `skip_layer_guidance_scale`, `skip_layer_guidance_start`, and `skip_layer_guidance_stop` are
55
+ defaulted to the recommendations by StabilityAI for Stable Diffusion 3.5 Medium.
56
+
57
+ Args:
58
+ guidance_scale (`float`, defaults to `7.5`):
59
+ The scale parameter for classifier-free guidance. Higher values result in stronger conditioning on the text
60
+ prompt, while lower values allow for more freedom in generation. Higher values may lead to saturation and
61
+ deterioration of image quality.
62
+ skip_layer_guidance_scale (`float`, defaults to `2.8`):
63
+ The scale parameter for skip layer guidance. Anatomy and structure coherence may improve with higher
64
+ values, but it may also lead to overexposure and saturation.
65
+ skip_layer_guidance_start (`float`, defaults to `0.01`):
66
+ The fraction of the total number of denoising steps after which skip layer guidance starts.
67
+ skip_layer_guidance_stop (`float`, defaults to `0.2`):
68
+ The fraction of the total number of denoising steps after which skip layer guidance stops.
69
+ skip_layer_guidance_layers (`int` or `list[int]`, *optional*):
70
+ The layer indices to apply skip layer guidance to. Can be a single integer or a list of integers. If not
71
+ provided, `skip_layer_config` must be provided. The recommended values are `[7, 8, 9]` for Stable Diffusion
72
+ 3.5 Medium.
73
+ skip_layer_config (`LayerSkipConfig` or `list[LayerSkipConfig]`, *optional*):
74
+ The configuration for the skip layer guidance. Can be a single `LayerSkipConfig` or a list of
75
+ `LayerSkipConfig`. If not provided, `skip_layer_guidance_layers` must be provided.
76
+ guidance_rescale (`float`, defaults to `0.0`):
77
+ The rescale factor applied to the noise predictions. This is used to improve image quality and fix
78
+ overexposure. Based on Section 3.4 from [Common Diffusion Noise Schedules and Sample Steps are
79
+ Flawed](https://huggingface.co/papers/2305.08891).
80
+ use_original_formulation (`bool`, defaults to `False`):
81
+ Whether to use the original formulation of classifier-free guidance as proposed in the paper. By default,
82
+ we use the diffusers-native implementation that has been in the codebase for a long time. See
83
+ [~guiders.classifier_free_guidance.ClassifierFreeGuidance] for more details.
84
+ start (`float`, defaults to `0.01`):
85
+ The fraction of the total number of denoising steps after which guidance starts.
86
+ stop (`float`, defaults to `0.2`):
87
+ The fraction of the total number of denoising steps after which guidance stops.
88
+ """
89
+
90
+ _input_predictions = ["pred_cond", "pred_uncond", "pred_cond_skip"]
91
+
92
+ @register_to_config
93
+ def __init__(
94
+ self,
95
+ guidance_scale: float = 7.5,
96
+ skip_layer_guidance_scale: float = 2.8,
97
+ skip_layer_guidance_start: float = 0.01,
98
+ skip_layer_guidance_stop: float = 0.2,
99
+ skip_layer_guidance_layers: int | list[int] | None = None,
100
+ skip_layer_config: LayerSkipConfig | list[LayerSkipConfig] | dict[str, Any] = None,
101
+ guidance_rescale: float = 0.0,
102
+ use_original_formulation: bool = False,
103
+ start: float = 0.0,
104
+ stop: float = 1.0,
105
+ enabled: bool = True,
106
+ ):
107
+ super().__init__(start, stop, enabled)
108
+
109
+ self.guidance_scale = guidance_scale
110
+ self.skip_layer_guidance_scale = skip_layer_guidance_scale
111
+ self.skip_layer_guidance_start = skip_layer_guidance_start
112
+ self.skip_layer_guidance_stop = skip_layer_guidance_stop
113
+ self.guidance_rescale = guidance_rescale
114
+ self.use_original_formulation = use_original_formulation
115
+
116
+ if not (0.0 <= skip_layer_guidance_start < 1.0):
117
+ raise ValueError(
118
+ f"Expected `skip_layer_guidance_start` to be between 0.0 and 1.0, but got {skip_layer_guidance_start}."
119
+ )
120
+ if not (skip_layer_guidance_start <= skip_layer_guidance_stop <= 1.0):
121
+ raise ValueError(
122
+ f"Expected `skip_layer_guidance_stop` to be between 0.0 and 1.0, but got {skip_layer_guidance_stop}."
123
+ )
124
+
125
+ if skip_layer_guidance_layers is None and skip_layer_config is None:
126
+ raise ValueError(
127
+ "Either `skip_layer_guidance_layers` or `skip_layer_config` must be provided to enable Skip Layer Guidance."
128
+ )
129
+ if skip_layer_guidance_layers is not None and skip_layer_config is not None:
130
+ raise ValueError("Only one of `skip_layer_guidance_layers` or `skip_layer_config` can be provided.")
131
+
132
+ if skip_layer_guidance_layers is not None:
133
+ if isinstance(skip_layer_guidance_layers, int):
134
+ skip_layer_guidance_layers = [skip_layer_guidance_layers]
135
+ if not isinstance(skip_layer_guidance_layers, list):
136
+ raise ValueError(
137
+ f"Expected `skip_layer_guidance_layers` to be an int or a list of ints, but got {type(skip_layer_guidance_layers)}."
138
+ )
139
+ skip_layer_config = [LayerSkipConfig(layer, fqn="auto") for layer in skip_layer_guidance_layers]
140
+
141
+ if isinstance(skip_layer_config, dict):
142
+ skip_layer_config = LayerSkipConfig.from_dict(skip_layer_config)
143
+
144
+ if isinstance(skip_layer_config, LayerSkipConfig):
145
+ skip_layer_config = [skip_layer_config]
146
+
147
+ if not isinstance(skip_layer_config, list):
148
+ raise ValueError(
149
+ f"Expected `skip_layer_config` to be a LayerSkipConfig or a list of LayerSkipConfig, but got {type(skip_layer_config)}."
150
+ )
151
+ elif isinstance(next(iter(skip_layer_config), None), dict):
152
+ skip_layer_config = [LayerSkipConfig.from_dict(config) for config in skip_layer_config]
153
+
154
+ self.skip_layer_config = skip_layer_config
155
+ self._skip_layer_hook_names = [f"SkipLayerGuidance_{i}" for i in range(len(self.skip_layer_config))]
156
+
157
+ def prepare_models(self, denoiser: torch.nn.Module) -> None:
158
+ self._count_prepared += 1
159
+ if self._is_slg_enabled() and self.is_conditional and self._count_prepared > 1:
160
+ for name, config in zip(self._skip_layer_hook_names, self.skip_layer_config):
161
+ _apply_layer_skip_hook(denoiser, config, name=name)
162
+
163
+ def cleanup_models(self, denoiser: torch.nn.Module) -> None:
164
+ if self._is_slg_enabled() and self.is_conditional and self._count_prepared > 1:
165
+ registry = HookRegistry.check_if_exists_or_initialize(denoiser)
166
+ # Remove the hooks after inference
167
+ for hook_name in self._skip_layer_hook_names:
168
+ registry.remove_hook(hook_name, recurse=True)
169
+
170
+ def prepare_inputs(self, data: dict[str, tuple[torch.Tensor, torch.Tensor]]) -> list["BlockState"]:
171
+ if self.num_conditions == 1:
172
+ tuple_indices = [0]
173
+ input_predictions = ["pred_cond"]
174
+ elif self.num_conditions == 2:
175
+ tuple_indices = [0, 1]
176
+ input_predictions = (
177
+ ["pred_cond", "pred_uncond"] if self._is_cfg_enabled() else ["pred_cond", "pred_cond_skip"]
178
+ )
179
+ else:
180
+ tuple_indices = [0, 1, 0]
181
+ input_predictions = ["pred_cond", "pred_uncond", "pred_cond_skip"]
182
+ data_batches = []
183
+ for tuple_idx, input_prediction in zip(tuple_indices, input_predictions):
184
+ data_batch = self._prepare_batch(data, tuple_idx, input_prediction)
185
+ data_batches.append(data_batch)
186
+ return data_batches
187
+
188
+ def prepare_inputs_from_block_state(
189
+ self, data: "BlockState", input_fields: dict[str, str | tuple[str, str]]
190
+ ) -> list["BlockState"]:
191
+ if self.num_conditions == 1:
192
+ tuple_indices = [0]
193
+ input_predictions = ["pred_cond"]
194
+ elif self.num_conditions == 2:
195
+ tuple_indices = [0, 1]
196
+ input_predictions = (
197
+ ["pred_cond", "pred_uncond"] if self._is_cfg_enabled() else ["pred_cond", "pred_cond_skip"]
198
+ )
199
+ else:
200
+ tuple_indices = [0, 1, 0]
201
+ input_predictions = ["pred_cond", "pred_uncond", "pred_cond_skip"]
202
+ data_batches = []
203
+ for tuple_idx, input_prediction in zip(tuple_indices, input_predictions):
204
+ data_batch = self._prepare_batch_from_block_state(input_fields, data, tuple_idx, input_prediction)
205
+ data_batches.append(data_batch)
206
+ return data_batches
207
+
208
+ def forward(
209
+ self,
210
+ pred_cond: torch.Tensor,
211
+ pred_uncond: torch.Tensor | None = None,
212
+ pred_cond_skip: torch.Tensor | None = None,
213
+ ) -> GuiderOutput:
214
+ pred = None
215
+
216
+ if not self._is_cfg_enabled() and not self._is_slg_enabled():
217
+ pred = pred_cond
218
+ elif not self._is_cfg_enabled():
219
+ shift = pred_cond - pred_cond_skip
220
+ pred = pred_cond if self.use_original_formulation else pred_cond_skip
221
+ pred = pred + self.skip_layer_guidance_scale * shift
222
+ elif not self._is_slg_enabled():
223
+ shift = pred_cond - pred_uncond
224
+ pred = pred_cond if self.use_original_formulation else pred_uncond
225
+ pred = pred + self.guidance_scale * shift
226
+ else:
227
+ shift = pred_cond - pred_uncond
228
+ shift_skip = pred_cond - pred_cond_skip
229
+ pred = pred_cond if self.use_original_formulation else pred_uncond
230
+ pred = pred + self.guidance_scale * shift + self.skip_layer_guidance_scale * shift_skip
231
+
232
+ if self.guidance_rescale > 0.0:
233
+ pred = rescale_noise_cfg(pred, pred_cond, self.guidance_rescale)
234
+
235
+ return GuiderOutput(pred=pred, pred_cond=pred_cond, pred_uncond=pred_uncond)
236
+
237
+ @property
238
+ def is_conditional(self) -> bool:
239
+ return self._count_prepared == 1 or self._count_prepared == 3
240
+
241
+ @property
242
+ def num_conditions(self) -> int:
243
+ num_conditions = 1
244
+ if self._is_cfg_enabled():
245
+ num_conditions += 1
246
+ if self._is_slg_enabled():
247
+ num_conditions += 1
248
+ return num_conditions
249
+
250
+ def _is_cfg_enabled(self) -> bool:
251
+ if not self._enabled:
252
+ return False
253
+
254
+ is_within_range = True
255
+ if self._num_inference_steps is not None:
256
+ skip_start_step = int(self._start * self._num_inference_steps)
257
+ skip_stop_step = int(self._stop * self._num_inference_steps)
258
+ is_within_range = skip_start_step <= self._step < skip_stop_step
259
+
260
+ is_close = False
261
+ if self.use_original_formulation:
262
+ is_close = math.isclose(self.guidance_scale, 0.0)
263
+ else:
264
+ is_close = math.isclose(self.guidance_scale, 1.0)
265
+
266
+ return is_within_range and not is_close
267
+
268
+ def _is_slg_enabled(self) -> bool:
269
+ if not self._enabled:
270
+ return False
271
+
272
+ is_within_range = True
273
+ if self._num_inference_steps is not None:
274
+ skip_start_step = int(self.skip_layer_guidance_start * self._num_inference_steps)
275
+ skip_stop_step = int(self.skip_layer_guidance_stop * self._num_inference_steps)
276
+ is_within_range = skip_start_step < self._step < skip_stop_step
277
+
278
+ is_zero = math.isclose(self.skip_layer_guidance_scale, 0.0)
279
+
280
+ return is_within_range and not is_zero
diffusers/guiders/smoothed_energy_guidance.py ADDED
@@ -0,0 +1,269 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 The HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ from __future__ import annotations
16
+
17
+ import math
18
+ from typing import TYPE_CHECKING
19
+
20
+ import torch
21
+
22
+ from ..configuration_utils import register_to_config
23
+ from ..hooks import HookRegistry
24
+ from ..hooks.smoothed_energy_guidance_utils import SmoothedEnergyGuidanceConfig, _apply_smoothed_energy_guidance_hook
25
+ from .guider_utils import BaseGuidance, GuiderOutput, rescale_noise_cfg
26
+
27
+
28
+ if TYPE_CHECKING:
29
+ from ..modular_pipelines.modular_pipeline import BlockState
30
+
31
+
32
+ class SmoothedEnergyGuidance(BaseGuidance):
33
+ """
34
+ Smoothed Energy Guidance (SEG): https://huggingface.co/papers/2408.00760
35
+
36
+ SEG is only supported as an experimental prototype feature for now, so the implementation may be modified in the
37
+ future without warning or guarantee of reproducibility. This implementation assumes:
38
+ - Generated images are square (height == width)
39
+ - The model does not combine different modalities together (e.g., text and image latent streams are not combined
40
+ together such as Flux)
41
+
42
+ Args:
43
+ guidance_scale (`float`, defaults to `7.5`):
44
+ The scale parameter for classifier-free guidance. Higher values result in stronger conditioning on the text
45
+ prompt, while lower values allow for more freedom in generation. Higher values may lead to saturation and
46
+ deterioration of image quality.
47
+ seg_guidance_scale (`float`, defaults to `3.0`):
48
+ The scale parameter for smoothed energy guidance. Anatomy and structure coherence may improve with higher
49
+ values, but it may also lead to overexposure and saturation.
50
+ seg_blur_sigma (`float`, defaults to `9999999.0`):
51
+ The amount by which we blur the attention weights. Setting this value greater than 9999.0 results in
52
+ infinite blur, which means uniform queries. Controlling it exponentially is empirically effective.
53
+ seg_blur_threshold_inf (`float`, defaults to `9999.0`):
54
+ The threshold above which the blur is considered infinite.
55
+ seg_guidance_start (`float`, defaults to `0.0`):
56
+ The fraction of the total number of denoising steps after which smoothed energy guidance starts.
57
+ seg_guidance_stop (`float`, defaults to `1.0`):
58
+ The fraction of the total number of denoising steps after which smoothed energy guidance stops.
59
+ seg_guidance_layers (`int` or `list[int]`, *optional*):
60
+ The layer indices to apply smoothed energy guidance to. Can be a single integer or a list of integers. If
61
+ not provided, `seg_guidance_config` must be provided. The recommended values are `[7, 8, 9]` for Stable
62
+ Diffusion 3.5 Medium.
63
+ seg_guidance_config (`SmoothedEnergyGuidanceConfig` or `list[SmoothedEnergyGuidanceConfig]`, *optional*):
64
+ The configuration for the smoothed energy layer guidance. Can be a single `SmoothedEnergyGuidanceConfig` or
65
+ a list of `SmoothedEnergyGuidanceConfig`. If not provided, `seg_guidance_layers` must be provided.
66
+ guidance_rescale (`float`, defaults to `0.0`):
67
+ The rescale factor applied to the noise predictions. This is used to improve image quality and fix
68
+ overexposure. Based on Section 3.4 from [Common Diffusion Noise Schedules and Sample Steps are
69
+ Flawed](https://huggingface.co/papers/2305.08891).
70
+ use_original_formulation (`bool`, defaults to `False`):
71
+ Whether to use the original formulation of classifier-free guidance as proposed in the paper. By default,
72
+ we use the diffusers-native implementation that has been in the codebase for a long time. See
73
+ [~guiders.classifier_free_guidance.ClassifierFreeGuidance] for more details.
74
+ start (`float`, defaults to `0.01`):
75
+ The fraction of the total number of denoising steps after which guidance starts.
76
+ stop (`float`, defaults to `0.2`):
77
+ The fraction of the total number of denoising steps after which guidance stops.
78
+ """
79
+
80
+ _input_predictions = ["pred_cond", "pred_uncond", "pred_cond_seg"]
81
+
82
+ @register_to_config
83
+ def __init__(
84
+ self,
85
+ guidance_scale: float = 7.5,
86
+ seg_guidance_scale: float = 2.8,
87
+ seg_blur_sigma: float = 9999999.0,
88
+ seg_blur_threshold_inf: float = 9999.0,
89
+ seg_guidance_start: float = 0.0,
90
+ seg_guidance_stop: float = 1.0,
91
+ seg_guidance_layers: int | list[int] | None = None,
92
+ seg_guidance_config: SmoothedEnergyGuidanceConfig | list[SmoothedEnergyGuidanceConfig] = None,
93
+ guidance_rescale: float = 0.0,
94
+ use_original_formulation: bool = False,
95
+ start: float = 0.0,
96
+ stop: float = 1.0,
97
+ enabled: bool = True,
98
+ ):
99
+ super().__init__(start, stop, enabled)
100
+
101
+ self.guidance_scale = guidance_scale
102
+ self.seg_guidance_scale = seg_guidance_scale
103
+ self.seg_blur_sigma = seg_blur_sigma
104
+ self.seg_blur_threshold_inf = seg_blur_threshold_inf
105
+ self.seg_guidance_start = seg_guidance_start
106
+ self.seg_guidance_stop = seg_guidance_stop
107
+ self.guidance_rescale = guidance_rescale
108
+ self.use_original_formulation = use_original_formulation
109
+
110
+ if not (0.0 <= seg_guidance_start < 1.0):
111
+ raise ValueError(f"Expected `seg_guidance_start` to be between 0.0 and 1.0, but got {seg_guidance_start}.")
112
+ if not (seg_guidance_start <= seg_guidance_stop <= 1.0):
113
+ raise ValueError(f"Expected `seg_guidance_stop` to be between 0.0 and 1.0, but got {seg_guidance_stop}.")
114
+
115
+ if seg_guidance_layers is None and seg_guidance_config is None:
116
+ raise ValueError(
117
+ "Either `seg_guidance_layers` or `seg_guidance_config` must be provided to enable Smoothed Energy Guidance."
118
+ )
119
+ if seg_guidance_layers is not None and seg_guidance_config is not None:
120
+ raise ValueError("Only one of `seg_guidance_layers` or `seg_guidance_config` can be provided.")
121
+
122
+ if seg_guidance_layers is not None:
123
+ if isinstance(seg_guidance_layers, int):
124
+ seg_guidance_layers = [seg_guidance_layers]
125
+ if not isinstance(seg_guidance_layers, list):
126
+ raise ValueError(
127
+ f"Expected `seg_guidance_layers` to be an int or a list of ints, but got {type(seg_guidance_layers)}."
128
+ )
129
+ seg_guidance_config = [SmoothedEnergyGuidanceConfig(layer, fqn="auto") for layer in seg_guidance_layers]
130
+
131
+ if isinstance(seg_guidance_config, dict):
132
+ seg_guidance_config = SmoothedEnergyGuidanceConfig.from_dict(seg_guidance_config)
133
+
134
+ if isinstance(seg_guidance_config, SmoothedEnergyGuidanceConfig):
135
+ seg_guidance_config = [seg_guidance_config]
136
+
137
+ if not isinstance(seg_guidance_config, list):
138
+ raise ValueError(
139
+ f"Expected `seg_guidance_config` to be a SmoothedEnergyGuidanceConfig or a list of SmoothedEnergyGuidanceConfig, but got {type(seg_guidance_config)}."
140
+ )
141
+ elif isinstance(next(iter(seg_guidance_config), None), dict):
142
+ seg_guidance_config = [SmoothedEnergyGuidanceConfig.from_dict(config) for config in seg_guidance_config]
143
+
144
+ self.seg_guidance_config = seg_guidance_config
145
+ self._seg_layer_hook_names = [f"SmoothedEnergyGuidance_{i}" for i in range(len(self.seg_guidance_config))]
146
+
147
+ def prepare_models(self, denoiser: torch.nn.Module) -> None:
148
+ if self._is_seg_enabled() and self.is_conditional and self._count_prepared > 1:
149
+ for name, config in zip(self._seg_layer_hook_names, self.seg_guidance_config):
150
+ _apply_smoothed_energy_guidance_hook(denoiser, config, self.seg_blur_sigma, name=name)
151
+
152
+ def cleanup_models(self, denoiser: torch.nn.Module):
153
+ if self._is_seg_enabled() and self.is_conditional and self._count_prepared > 1:
154
+ registry = HookRegistry.check_if_exists_or_initialize(denoiser)
155
+ # Remove the hooks after inference
156
+ for hook_name in self._seg_layer_hook_names:
157
+ registry.remove_hook(hook_name, recurse=True)
158
+
159
+ def prepare_inputs(self, data: dict[str, tuple[torch.Tensor, torch.Tensor]]) -> list["BlockState"]:
160
+ if self.num_conditions == 1:
161
+ tuple_indices = [0]
162
+ input_predictions = ["pred_cond"]
163
+ elif self.num_conditions == 2:
164
+ tuple_indices = [0, 1]
165
+ input_predictions = (
166
+ ["pred_cond", "pred_uncond"] if self._is_cfg_enabled() else ["pred_cond", "pred_cond_seg"]
167
+ )
168
+ else:
169
+ tuple_indices = [0, 1, 0]
170
+ input_predictions = ["pred_cond", "pred_uncond", "pred_cond_seg"]
171
+ data_batches = []
172
+ for tuple_idx, input_prediction in zip(tuple_indices, input_predictions):
173
+ data_batch = self._prepare_batch(data, tuple_idx, input_prediction)
174
+ data_batches.append(data_batch)
175
+ return data_batches
176
+
177
+ def prepare_inputs_from_block_state(
178
+ self, data: "BlockState", input_fields: dict[str, str | tuple[str, str]]
179
+ ) -> list["BlockState"]:
180
+ if self.num_conditions == 1:
181
+ tuple_indices = [0]
182
+ input_predictions = ["pred_cond"]
183
+ elif self.num_conditions == 2:
184
+ tuple_indices = [0, 1]
185
+ input_predictions = (
186
+ ["pred_cond", "pred_uncond"] if self._is_cfg_enabled() else ["pred_cond", "pred_cond_seg"]
187
+ )
188
+ else:
189
+ tuple_indices = [0, 1, 0]
190
+ input_predictions = ["pred_cond", "pred_uncond", "pred_cond_seg"]
191
+ data_batches = []
192
+ for tuple_idx, input_prediction in zip(tuple_indices, input_predictions):
193
+ data_batch = self._prepare_batch_from_block_state(input_fields, data, tuple_idx, input_prediction)
194
+ data_batches.append(data_batch)
195
+ return data_batches
196
+
197
+ def forward(
198
+ self,
199
+ pred_cond: torch.Tensor,
200
+ pred_uncond: torch.Tensor | None = None,
201
+ pred_cond_seg: torch.Tensor | None = None,
202
+ ) -> GuiderOutput:
203
+ pred = None
204
+
205
+ if not self._is_cfg_enabled() and not self._is_seg_enabled():
206
+ pred = pred_cond
207
+ elif not self._is_cfg_enabled():
208
+ shift = pred_cond - pred_cond_seg
209
+ pred = pred_cond if self.use_original_formulation else pred_cond_seg
210
+ pred = pred + self.seg_guidance_scale * shift
211
+ elif not self._is_seg_enabled():
212
+ shift = pred_cond - pred_uncond
213
+ pred = pred_cond if self.use_original_formulation else pred_uncond
214
+ pred = pred + self.guidance_scale * shift
215
+ else:
216
+ shift = pred_cond - pred_uncond
217
+ shift_seg = pred_cond - pred_cond_seg
218
+ pred = pred_cond if self.use_original_formulation else pred_uncond
219
+ pred = pred + self.guidance_scale * shift + self.seg_guidance_scale * shift_seg
220
+
221
+ if self.guidance_rescale > 0.0:
222
+ pred = rescale_noise_cfg(pred, pred_cond, self.guidance_rescale)
223
+
224
+ return GuiderOutput(pred=pred, pred_cond=pred_cond, pred_uncond=pred_uncond)
225
+
226
+ @property
227
+ def is_conditional(self) -> bool:
228
+ return self._count_prepared == 1 or self._count_prepared == 3
229
+
230
+ @property
231
+ def num_conditions(self) -> int:
232
+ num_conditions = 1
233
+ if self._is_cfg_enabled():
234
+ num_conditions += 1
235
+ if self._is_seg_enabled():
236
+ num_conditions += 1
237
+ return num_conditions
238
+
239
+ def _is_cfg_enabled(self) -> bool:
240
+ if not self._enabled:
241
+ return False
242
+
243
+ is_within_range = True
244
+ if self._num_inference_steps is not None:
245
+ skip_start_step = int(self._start * self._num_inference_steps)
246
+ skip_stop_step = int(self._stop * self._num_inference_steps)
247
+ is_within_range = skip_start_step <= self._step < skip_stop_step
248
+
249
+ is_close = False
250
+ if self.use_original_formulation:
251
+ is_close = math.isclose(self.guidance_scale, 0.0)
252
+ else:
253
+ is_close = math.isclose(self.guidance_scale, 1.0)
254
+
255
+ return is_within_range and not is_close
256
+
257
+ def _is_seg_enabled(self) -> bool:
258
+ if not self._enabled:
259
+ return False
260
+
261
+ is_within_range = True
262
+ if self._num_inference_steps is not None:
263
+ skip_start_step = int(self.seg_guidance_start * self._num_inference_steps)
264
+ skip_stop_step = int(self.seg_guidance_stop * self._num_inference_steps)
265
+ is_within_range = skip_start_step < self._step < skip_stop_step
266
+
267
+ is_zero = math.isclose(self.seg_guidance_scale, 0.0)
268
+
269
+ return is_within_range and not is_zero
diffusers/guiders/tangential_classifier_free_guidance.py ADDED
@@ -0,0 +1,151 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 The HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ from __future__ import annotations
16
+
17
+ import math
18
+ from typing import TYPE_CHECKING
19
+
20
+ import torch
21
+
22
+ from ..configuration_utils import register_to_config
23
+ from .guider_utils import BaseGuidance, GuiderOutput, rescale_noise_cfg
24
+
25
+
26
+ if TYPE_CHECKING:
27
+ from ..modular_pipelines.modular_pipeline import BlockState
28
+
29
+
30
+ class TangentialClassifierFreeGuidance(BaseGuidance):
31
+ """
32
+ Tangential Classifier Free Guidance (TCFG): https://huggingface.co/papers/2503.18137
33
+
34
+ Args:
35
+ guidance_scale (`float`, defaults to `7.5`):
36
+ The scale parameter for classifier-free guidance. Higher values result in stronger conditioning on the text
37
+ prompt, while lower values allow for more freedom in generation. Higher values may lead to saturation and
38
+ deterioration of image quality.
39
+ guidance_rescale (`float`, defaults to `0.0`):
40
+ The rescale factor applied to the noise predictions. This is used to improve image quality and fix
41
+ overexposure. Based on Section 3.4 from [Common Diffusion Noise Schedules and Sample Steps are
42
+ Flawed](https://huggingface.co/papers/2305.08891).
43
+ use_original_formulation (`bool`, defaults to `False`):
44
+ Whether to use the original formulation of classifier-free guidance as proposed in the paper. By default,
45
+ we use the diffusers-native implementation that has been in the codebase for a long time. See
46
+ [~guiders.classifier_free_guidance.ClassifierFreeGuidance] for more details.
47
+ start (`float`, defaults to `0.0`):
48
+ The fraction of the total number of denoising steps after which guidance starts.
49
+ stop (`float`, defaults to `1.0`):
50
+ The fraction of the total number of denoising steps after which guidance stops.
51
+ """
52
+
53
+ _input_predictions = ["pred_cond", "pred_uncond"]
54
+
55
+ @register_to_config
56
+ def __init__(
57
+ self,
58
+ guidance_scale: float = 7.5,
59
+ guidance_rescale: float = 0.0,
60
+ use_original_formulation: bool = False,
61
+ start: float = 0.0,
62
+ stop: float = 1.0,
63
+ enabled: bool = True,
64
+ ):
65
+ super().__init__(start, stop, enabled)
66
+
67
+ self.guidance_scale = guidance_scale
68
+ self.guidance_rescale = guidance_rescale
69
+ self.use_original_formulation = use_original_formulation
70
+
71
+ def prepare_inputs(self, data: dict[str, tuple[torch.Tensor, torch.Tensor]]) -> list["BlockState"]:
72
+ tuple_indices = [0] if self.num_conditions == 1 else [0, 1]
73
+ data_batches = []
74
+ for tuple_idx, input_prediction in zip(tuple_indices, self._input_predictions):
75
+ data_batch = self._prepare_batch(data, tuple_idx, input_prediction)
76
+ data_batches.append(data_batch)
77
+ return data_batches
78
+
79
+ def prepare_inputs_from_block_state(
80
+ self, data: "BlockState", input_fields: dict[str, str | tuple[str, str]]
81
+ ) -> list["BlockState"]:
82
+ tuple_indices = [0] if self.num_conditions == 1 else [0, 1]
83
+ data_batches = []
84
+ for tuple_idx, input_prediction in zip(tuple_indices, self._input_predictions):
85
+ data_batch = self._prepare_batch_from_block_state(input_fields, data, tuple_idx, input_prediction)
86
+ data_batches.append(data_batch)
87
+ return data_batches
88
+
89
+ def forward(self, pred_cond: torch.Tensor, pred_uncond: torch.Tensor | None = None) -> GuiderOutput:
90
+ pred = None
91
+
92
+ if not self._is_tcfg_enabled():
93
+ pred = pred_cond
94
+ else:
95
+ pred = normalized_guidance(pred_cond, pred_uncond, self.guidance_scale, self.use_original_formulation)
96
+
97
+ if self.guidance_rescale > 0.0:
98
+ pred = rescale_noise_cfg(pred, pred_cond, self.guidance_rescale)
99
+
100
+ return GuiderOutput(pred=pred, pred_cond=pred_cond, pred_uncond=pred_uncond)
101
+
102
+ @property
103
+ def is_conditional(self) -> bool:
104
+ return self._num_outputs_prepared == 1
105
+
106
+ @property
107
+ def num_conditions(self) -> int:
108
+ num_conditions = 1
109
+ if self._is_tcfg_enabled():
110
+ num_conditions += 1
111
+ return num_conditions
112
+
113
+ def _is_tcfg_enabled(self) -> bool:
114
+ if not self._enabled:
115
+ return False
116
+
117
+ is_within_range = True
118
+ if self._num_inference_steps is not None:
119
+ skip_start_step = int(self._start * self._num_inference_steps)
120
+ skip_stop_step = int(self._stop * self._num_inference_steps)
121
+ is_within_range = skip_start_step <= self._step < skip_stop_step
122
+
123
+ is_close = False
124
+ if self.use_original_formulation:
125
+ is_close = math.isclose(self.guidance_scale, 0.0)
126
+ else:
127
+ is_close = math.isclose(self.guidance_scale, 1.0)
128
+
129
+ return is_within_range and not is_close
130
+
131
+
132
+ def normalized_guidance(
133
+ pred_cond: torch.Tensor, pred_uncond: torch.Tensor, guidance_scale: float, use_original_formulation: bool = False
134
+ ) -> torch.Tensor:
135
+ cond_dtype = pred_cond.dtype
136
+ preds = torch.stack([pred_cond, pred_uncond], dim=1).float()
137
+ preds = preds.flatten(2)
138
+ U, S, Vh = torch.linalg.svd(preds, full_matrices=False)
139
+ Vh_modified = Vh.clone()
140
+ Vh_modified[:, 1] = 0
141
+
142
+ uncond_flat = pred_uncond.reshape(pred_uncond.size(0), 1, -1).float()
143
+ x_Vh = torch.matmul(uncond_flat, Vh.transpose(-2, -1))
144
+ x_Vh_V = torch.matmul(x_Vh, Vh_modified)
145
+ pred_uncond = x_Vh_V.reshape(pred_uncond.shape).to(cond_dtype)
146
+
147
+ pred = pred_cond if use_original_formulation else pred_uncond
148
+ shift = pred_cond - pred_uncond
149
+ pred = pred + guidance_scale * shift
150
+
151
+ return pred
diffusers/hooks/__init__.py ADDED
@@ -0,0 +1,30 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2024 The HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ from ..utils import is_torch_available
16
+
17
+
18
+ if is_torch_available():
19
+ from .context_parallel import apply_context_parallel
20
+ from .faster_cache import FasterCacheConfig, apply_faster_cache
21
+ from .first_block_cache import FirstBlockCacheConfig, apply_first_block_cache
22
+ from .group_offloading import apply_group_offloading
23
+ from .hooks import HookRegistry, ModelHook
24
+ from .layer_skip import LayerSkipConfig, apply_layer_skip
25
+ from .layerwise_casting import apply_layerwise_casting, apply_layerwise_casting_hook
26
+ from .mag_cache import MagCacheConfig, apply_mag_cache
27
+ from .pyramid_attention_broadcast import PyramidAttentionBroadcastConfig, apply_pyramid_attention_broadcast
28
+ from .smoothed_energy_guidance_utils import SmoothedEnergyGuidanceConfig
29
+ from .taylorseer_cache import TaylorSeerCacheConfig, apply_taylorseer_cache
30
+ from .text_kv_cache import TextKVCacheConfig, apply_text_kv_cache
diffusers/hooks/_common.py ADDED
@@ -0,0 +1,61 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 The HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ import torch
16
+
17
+ from ..models.attention import AttentionModuleMixin, FeedForward, LuminaFeedForward
18
+ from ..models.attention_processor import Attention, MochiAttention
19
+
20
+
21
+ _ATTENTION_CLASSES = (Attention, MochiAttention, AttentionModuleMixin)
22
+ _FEEDFORWARD_CLASSES = (FeedForward, LuminaFeedForward)
23
+
24
+ _SPATIAL_TRANSFORMER_BLOCK_IDENTIFIERS = (
25
+ "blocks",
26
+ "transformer_blocks",
27
+ "single_transformer_blocks",
28
+ "layers",
29
+ "visual_transformer_blocks",
30
+ )
31
+ _TEMPORAL_TRANSFORMER_BLOCK_IDENTIFIERS = ("temporal_transformer_blocks",)
32
+ _CROSS_TRANSFORMER_BLOCK_IDENTIFIERS = ("blocks", "transformer_blocks", "layers")
33
+
34
+ _ALL_TRANSFORMER_BLOCK_IDENTIFIERS = tuple(
35
+ {
36
+ *_SPATIAL_TRANSFORMER_BLOCK_IDENTIFIERS,
37
+ *_TEMPORAL_TRANSFORMER_BLOCK_IDENTIFIERS,
38
+ *_CROSS_TRANSFORMER_BLOCK_IDENTIFIERS,
39
+ }
40
+ )
41
+
42
+ # Layers supported for group offloading and layerwise casting
43
+ _GO_LC_SUPPORTED_PYTORCH_LAYERS = (
44
+ torch.nn.Conv1d,
45
+ torch.nn.Conv2d,
46
+ torch.nn.Conv3d,
47
+ torch.nn.ConvTranspose1d,
48
+ torch.nn.ConvTranspose2d,
49
+ torch.nn.ConvTranspose3d,
50
+ torch.nn.Linear,
51
+ torch.nn.Embedding,
52
+ # TODO(aryan): look into torch.nn.LayerNorm, torch.nn.GroupNorm later, seems to be causing some issues with CogVideoX
53
+ # because of double invocation of the same norm layer in CogVideoXLayerNorm
54
+ )
55
+
56
+
57
+ def _get_submodule_from_fqn(module: torch.nn.Module, fqn: str) -> torch.nn.Module | None:
58
+ for submodule_name, submodule in module.named_modules():
59
+ if submodule_name == fqn:
60
+ return submodule
61
+ return None
diffusers/hooks/_helpers.py ADDED
@@ -0,0 +1,401 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 The HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ import inspect
16
+ from dataclasses import dataclass
17
+ from typing import Any, Callable, Type
18
+
19
+
20
+ @dataclass
21
+ class AttentionProcessorMetadata:
22
+ skip_processor_output_fn: Callable[[Any], Any]
23
+
24
+
25
+ @dataclass
26
+ class TransformerBlockMetadata:
27
+ return_hidden_states_index: int = None
28
+ return_encoder_hidden_states_index: int = None
29
+ hidden_states_argument_name: str = "hidden_states"
30
+
31
+ _cls: Type = None
32
+ _cached_parameter_indices: dict[str, int] = None
33
+
34
+ def _get_parameter_from_args_kwargs(self, identifier: str, args=(), kwargs=None):
35
+ kwargs = kwargs or {}
36
+ if identifier in kwargs:
37
+ return kwargs[identifier]
38
+ if self._cached_parameter_indices is not None:
39
+ return args[self._cached_parameter_indices[identifier]]
40
+ if self._cls is None:
41
+ raise ValueError("Model class is not set for metadata.")
42
+ parameters = list(inspect.signature(self._cls.forward).parameters.keys())
43
+ parameters = parameters[1:] # skip `self`
44
+ self._cached_parameter_indices = {param: i for i, param in enumerate(parameters)}
45
+ if identifier not in self._cached_parameter_indices:
46
+ raise ValueError(f"Parameter '{identifier}' not found in function signature but was requested.")
47
+ index = self._cached_parameter_indices[identifier]
48
+ if index >= len(args):
49
+ raise ValueError(f"Expected {index} arguments but got {len(args)}.")
50
+ return args[index]
51
+
52
+
53
+ class AttentionProcessorRegistry:
54
+ _registry = {}
55
+ # TODO(aryan): this is only required for the time being because we need to do the registrations
56
+ # for classes. If we do it eagerly, i.e. call the functions in global scope, we will get circular
57
+ # import errors because of the models imported in this file.
58
+ _is_registered = False
59
+
60
+ @classmethod
61
+ def register(cls, model_class: Type, metadata: AttentionProcessorMetadata):
62
+ cls._register()
63
+ cls._registry[model_class] = metadata
64
+
65
+ @classmethod
66
+ def get(cls, model_class: Type) -> AttentionProcessorMetadata:
67
+ cls._register()
68
+ if model_class not in cls._registry:
69
+ raise ValueError(f"Model class {model_class} not registered.")
70
+ return cls._registry[model_class]
71
+
72
+ @classmethod
73
+ def _register(cls):
74
+ if cls._is_registered:
75
+ return
76
+ cls._is_registered = True
77
+ _register_attention_processors_metadata()
78
+
79
+
80
+ class TransformerBlockRegistry:
81
+ _registry = {}
82
+ # TODO(aryan): this is only required for the time being because we need to do the registrations
83
+ # for classes. If we do it eagerly, i.e. call the functions in global scope, we will get circular
84
+ # import errors because of the models imported in this file.
85
+ _is_registered = False
86
+
87
+ @classmethod
88
+ def register(cls, model_class: Type, metadata: TransformerBlockMetadata):
89
+ cls._register()
90
+ metadata._cls = model_class
91
+ cls._registry[model_class] = metadata
92
+
93
+ @classmethod
94
+ def get(cls, model_class: Type) -> TransformerBlockMetadata:
95
+ cls._register()
96
+ if model_class not in cls._registry:
97
+ raise ValueError(f"Model class {model_class} not registered.")
98
+ return cls._registry[model_class]
99
+
100
+ @classmethod
101
+ def _register(cls):
102
+ if cls._is_registered:
103
+ return
104
+ cls._is_registered = True
105
+ _register_transformer_blocks_metadata()
106
+
107
+
108
+ def _register_attention_processors_metadata():
109
+ from ..models.attention_processor import AttnProcessor2_0
110
+ from ..models.transformers.transformer_cogview4 import CogView4AttnProcessor
111
+ from ..models.transformers.transformer_flux import FluxAttnProcessor
112
+ from ..models.transformers.transformer_hunyuanimage import HunyuanImageAttnProcessor
113
+ from ..models.transformers.transformer_qwenimage import QwenDoubleStreamAttnProcessor2_0
114
+ from ..models.transformers.transformer_wan import WanAttnProcessor2_0
115
+ from ..models.transformers.transformer_z_image import ZSingleStreamAttnProcessor
116
+
117
+ # AttnProcessor2_0
118
+ AttentionProcessorRegistry.register(
119
+ model_class=AttnProcessor2_0,
120
+ metadata=AttentionProcessorMetadata(
121
+ skip_processor_output_fn=_skip_proc_output_fn_Attention_AttnProcessor2_0,
122
+ ),
123
+ )
124
+
125
+ # CogView4AttnProcessor
126
+ AttentionProcessorRegistry.register(
127
+ model_class=CogView4AttnProcessor,
128
+ metadata=AttentionProcessorMetadata(
129
+ skip_processor_output_fn=_skip_proc_output_fn_Attention_CogView4AttnProcessor,
130
+ ),
131
+ )
132
+
133
+ # WanAttnProcessor2_0
134
+ AttentionProcessorRegistry.register(
135
+ model_class=WanAttnProcessor2_0,
136
+ metadata=AttentionProcessorMetadata(
137
+ skip_processor_output_fn=_skip_proc_output_fn_Attention_WanAttnProcessor2_0,
138
+ ),
139
+ )
140
+
141
+ # FluxAttnProcessor
142
+ AttentionProcessorRegistry.register(
143
+ model_class=FluxAttnProcessor,
144
+ metadata=AttentionProcessorMetadata(skip_processor_output_fn=_skip_proc_output_fn_Attention_FluxAttnProcessor),
145
+ )
146
+
147
+ # QwenDoubleStreamAttnProcessor2
148
+ AttentionProcessorRegistry.register(
149
+ model_class=QwenDoubleStreamAttnProcessor2_0,
150
+ metadata=AttentionProcessorMetadata(
151
+ skip_processor_output_fn=_skip_proc_output_fn_Attention_QwenDoubleStreamAttnProcessor2_0
152
+ ),
153
+ )
154
+
155
+ # HunyuanImageAttnProcessor
156
+ AttentionProcessorRegistry.register(
157
+ model_class=HunyuanImageAttnProcessor,
158
+ metadata=AttentionProcessorMetadata(
159
+ skip_processor_output_fn=_skip_proc_output_fn_Attention_HunyuanImageAttnProcessor,
160
+ ),
161
+ )
162
+
163
+ # ZSingleStreamAttnProcessor
164
+ AttentionProcessorRegistry.register(
165
+ model_class=ZSingleStreamAttnProcessor,
166
+ metadata=AttentionProcessorMetadata(
167
+ skip_processor_output_fn=_skip_proc_output_fn_Attention_ZSingleStreamAttnProcessor,
168
+ ),
169
+ )
170
+
171
+
172
+ def _register_transformer_blocks_metadata():
173
+ from ..models.attention import BasicTransformerBlock, JointTransformerBlock
174
+ from ..models.transformers.cogvideox_transformer_3d import CogVideoXBlock
175
+ from ..models.transformers.transformer_bria import BriaTransformerBlock
176
+ from ..models.transformers.transformer_cogview4 import CogView4TransformerBlock
177
+ from ..models.transformers.transformer_flux import FluxSingleTransformerBlock, FluxTransformerBlock
178
+ from ..models.transformers.transformer_hunyuan_video import (
179
+ HunyuanVideoSingleTransformerBlock,
180
+ HunyuanVideoTokenReplaceSingleTransformerBlock,
181
+ HunyuanVideoTokenReplaceTransformerBlock,
182
+ HunyuanVideoTransformerBlock,
183
+ )
184
+ from ..models.transformers.transformer_hunyuanimage import (
185
+ HunyuanImageSingleTransformerBlock,
186
+ HunyuanImageTransformerBlock,
187
+ )
188
+ from ..models.transformers.transformer_kandinsky import Kandinsky5TransformerDecoderBlock
189
+ from ..models.transformers.transformer_ltx import LTXVideoTransformerBlock
190
+ from ..models.transformers.transformer_mochi import MochiTransformerBlock
191
+ from ..models.transformers.transformer_motif_video import (
192
+ MotifVideoSingleTransformerBlock,
193
+ MotifVideoTransformerBlock,
194
+ )
195
+ from ..models.transformers.transformer_qwenimage import QwenImageTransformerBlock
196
+ from ..models.transformers.transformer_wan import WanTransformerBlock
197
+ from ..models.transformers.transformer_z_image import ZImageTransformerBlock
198
+
199
+ # BasicTransformerBlock
200
+ TransformerBlockRegistry.register(
201
+ model_class=BasicTransformerBlock,
202
+ metadata=TransformerBlockMetadata(
203
+ return_hidden_states_index=0,
204
+ return_encoder_hidden_states_index=None,
205
+ ),
206
+ )
207
+ TransformerBlockRegistry.register(
208
+ model_class=BriaTransformerBlock,
209
+ metadata=TransformerBlockMetadata(
210
+ return_hidden_states_index=0,
211
+ return_encoder_hidden_states_index=None,
212
+ ),
213
+ )
214
+
215
+ # CogVideoX
216
+ TransformerBlockRegistry.register(
217
+ model_class=CogVideoXBlock,
218
+ metadata=TransformerBlockMetadata(
219
+ return_hidden_states_index=0,
220
+ return_encoder_hidden_states_index=1,
221
+ ),
222
+ )
223
+
224
+ # CogView4
225
+ TransformerBlockRegistry.register(
226
+ model_class=CogView4TransformerBlock,
227
+ metadata=TransformerBlockMetadata(
228
+ return_hidden_states_index=0,
229
+ return_encoder_hidden_states_index=1,
230
+ ),
231
+ )
232
+
233
+ # Flux
234
+ TransformerBlockRegistry.register(
235
+ model_class=FluxTransformerBlock,
236
+ metadata=TransformerBlockMetadata(
237
+ return_hidden_states_index=1,
238
+ return_encoder_hidden_states_index=0,
239
+ ),
240
+ )
241
+ TransformerBlockRegistry.register(
242
+ model_class=FluxSingleTransformerBlock,
243
+ metadata=TransformerBlockMetadata(
244
+ return_hidden_states_index=1,
245
+ return_encoder_hidden_states_index=0,
246
+ ),
247
+ )
248
+
249
+ # HunyuanVideo
250
+ TransformerBlockRegistry.register(
251
+ model_class=HunyuanVideoTransformerBlock,
252
+ metadata=TransformerBlockMetadata(
253
+ return_hidden_states_index=0,
254
+ return_encoder_hidden_states_index=1,
255
+ ),
256
+ )
257
+ TransformerBlockRegistry.register(
258
+ model_class=HunyuanVideoSingleTransformerBlock,
259
+ metadata=TransformerBlockMetadata(
260
+ return_hidden_states_index=0,
261
+ return_encoder_hidden_states_index=1,
262
+ ),
263
+ )
264
+ TransformerBlockRegistry.register(
265
+ model_class=HunyuanVideoTokenReplaceTransformerBlock,
266
+ metadata=TransformerBlockMetadata(
267
+ return_hidden_states_index=0,
268
+ return_encoder_hidden_states_index=1,
269
+ ),
270
+ )
271
+ TransformerBlockRegistry.register(
272
+ model_class=HunyuanVideoTokenReplaceSingleTransformerBlock,
273
+ metadata=TransformerBlockMetadata(
274
+ return_hidden_states_index=0,
275
+ return_encoder_hidden_states_index=1,
276
+ ),
277
+ )
278
+
279
+ # LTXVideo
280
+ TransformerBlockRegistry.register(
281
+ model_class=LTXVideoTransformerBlock,
282
+ metadata=TransformerBlockMetadata(
283
+ return_hidden_states_index=0,
284
+ return_encoder_hidden_states_index=None,
285
+ ),
286
+ )
287
+
288
+ # Mochi
289
+ TransformerBlockRegistry.register(
290
+ model_class=MochiTransformerBlock,
291
+ metadata=TransformerBlockMetadata(
292
+ return_hidden_states_index=0,
293
+ return_encoder_hidden_states_index=1,
294
+ ),
295
+ )
296
+
297
+ # MotifVideo
298
+ TransformerBlockRegistry.register(
299
+ model_class=MotifVideoTransformerBlock,
300
+ metadata=TransformerBlockMetadata(
301
+ return_hidden_states_index=0,
302
+ return_encoder_hidden_states_index=1,
303
+ ),
304
+ )
305
+ TransformerBlockRegistry.register(
306
+ model_class=MotifVideoSingleTransformerBlock,
307
+ metadata=TransformerBlockMetadata(
308
+ return_hidden_states_index=0,
309
+ return_encoder_hidden_states_index=1,
310
+ ),
311
+ )
312
+
313
+ # Wan
314
+ TransformerBlockRegistry.register(
315
+ model_class=WanTransformerBlock,
316
+ metadata=TransformerBlockMetadata(
317
+ return_hidden_states_index=0,
318
+ return_encoder_hidden_states_index=None,
319
+ ),
320
+ )
321
+
322
+ # QwenImage
323
+ TransformerBlockRegistry.register(
324
+ model_class=QwenImageTransformerBlock,
325
+ metadata=TransformerBlockMetadata(
326
+ return_hidden_states_index=1,
327
+ return_encoder_hidden_states_index=0,
328
+ ),
329
+ )
330
+
331
+ # HunyuanImage2.1
332
+ TransformerBlockRegistry.register(
333
+ model_class=HunyuanImageTransformerBlock,
334
+ metadata=TransformerBlockMetadata(
335
+ return_hidden_states_index=0,
336
+ return_encoder_hidden_states_index=1,
337
+ ),
338
+ )
339
+ TransformerBlockRegistry.register(
340
+ model_class=HunyuanImageSingleTransformerBlock,
341
+ metadata=TransformerBlockMetadata(
342
+ return_hidden_states_index=0,
343
+ return_encoder_hidden_states_index=1,
344
+ ),
345
+ )
346
+
347
+ # ZImage
348
+ TransformerBlockRegistry.register(
349
+ model_class=ZImageTransformerBlock,
350
+ metadata=TransformerBlockMetadata(
351
+ return_hidden_states_index=0,
352
+ return_encoder_hidden_states_index=None,
353
+ ),
354
+ )
355
+
356
+ TransformerBlockRegistry.register(
357
+ model_class=JointTransformerBlock,
358
+ metadata=TransformerBlockMetadata(
359
+ return_hidden_states_index=1,
360
+ return_encoder_hidden_states_index=0,
361
+ ),
362
+ )
363
+
364
+ # Kandinsky 5.0 (Kandinsky5TransformerDecoderBlock)
365
+ TransformerBlockRegistry.register(
366
+ model_class=Kandinsky5TransformerDecoderBlock,
367
+ metadata=TransformerBlockMetadata(
368
+ return_hidden_states_index=0,
369
+ return_encoder_hidden_states_index=None,
370
+ hidden_states_argument_name="visual_embed",
371
+ ),
372
+ )
373
+
374
+
375
+ # fmt: off
376
+ def _skip_attention___ret___hidden_states(self, *args, **kwargs):
377
+ hidden_states = kwargs.get("hidden_states", None)
378
+ if hidden_states is None and len(args) > 0:
379
+ hidden_states = args[0]
380
+ return hidden_states
381
+
382
+
383
+ def _skip_attention___ret___hidden_states___encoder_hidden_states(self, *args, **kwargs):
384
+ hidden_states = kwargs.get("hidden_states", None)
385
+ encoder_hidden_states = kwargs.get("encoder_hidden_states", None)
386
+ if hidden_states is None and len(args) > 0:
387
+ hidden_states = args[0]
388
+ if encoder_hidden_states is None and len(args) > 1:
389
+ encoder_hidden_states = args[1]
390
+ return hidden_states, encoder_hidden_states
391
+
392
+
393
+ _skip_proc_output_fn_Attention_AttnProcessor2_0 = _skip_attention___ret___hidden_states
394
+ _skip_proc_output_fn_Attention_CogView4AttnProcessor = _skip_attention___ret___hidden_states___encoder_hidden_states
395
+ _skip_proc_output_fn_Attention_WanAttnProcessor2_0 = _skip_attention___ret___hidden_states
396
+ # not sure what this is yet.
397
+ _skip_proc_output_fn_Attention_FluxAttnProcessor = _skip_attention___ret___hidden_states
398
+ _skip_proc_output_fn_Attention_QwenDoubleStreamAttnProcessor2_0 = _skip_attention___ret___hidden_states
399
+ _skip_proc_output_fn_Attention_HunyuanImageAttnProcessor = _skip_attention___ret___hidden_states
400
+ _skip_proc_output_fn_Attention_ZSingleStreamAttnProcessor = _skip_attention___ret___hidden_states
401
+ # fmt: on
diffusers/hooks/context_parallel.py ADDED
@@ -0,0 +1,382 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 The HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+ import copy
15
+ import inspect
16
+ from dataclasses import dataclass
17
+ from typing import Type
18
+
19
+ import torch
20
+ import torch.distributed as dist
21
+
22
+
23
+ if torch.distributed.is_available():
24
+ import torch.distributed._functional_collectives as funcol
25
+
26
+ from ..models._modeling_parallel import (
27
+ ContextParallelConfig,
28
+ ContextParallelInput,
29
+ ContextParallelModelPlan,
30
+ ContextParallelOutput,
31
+ gather_size_by_comm,
32
+ )
33
+ from ..utils import get_logger
34
+ from ..utils.torch_utils import lru_cache_unless_export, maybe_allow_in_graph, unwrap_module
35
+ from .hooks import HookRegistry, ModelHook
36
+
37
+
38
+ logger = get_logger(__name__) # pylint: disable=invalid-name
39
+
40
+ _CONTEXT_PARALLEL_INPUT_HOOK_TEMPLATE = "cp_input---{}"
41
+ _CONTEXT_PARALLEL_OUTPUT_HOOK_TEMPLATE = "cp_output---{}"
42
+
43
+
44
+ # TODO(aryan): consolidate with ._helpers.TransformerBlockMetadata
45
+ @dataclass
46
+ class ModuleForwardMetadata:
47
+ cached_parameter_indices: dict[str, int] = None
48
+ _cls: Type = None
49
+
50
+ def _get_parameter_from_args_kwargs(self, identifier: str, args=(), kwargs=None):
51
+ kwargs = kwargs or {}
52
+
53
+ if identifier in kwargs:
54
+ return kwargs[identifier], True, None
55
+
56
+ if self.cached_parameter_indices is not None:
57
+ index = self.cached_parameter_indices.get(identifier, None)
58
+ if index is None:
59
+ raise ValueError(f"Parameter '{identifier}' not found in cached indices.")
60
+ return args[index], False, index
61
+
62
+ if self._cls is None:
63
+ raise ValueError("Model class is not set for metadata.")
64
+
65
+ parameters = list(inspect.signature(self._cls.forward).parameters.keys())
66
+ parameters = parameters[1:] # skip `self`
67
+ self.cached_parameter_indices = {param: i for i, param in enumerate(parameters)}
68
+
69
+ if identifier not in self.cached_parameter_indices:
70
+ raise ValueError(f"Parameter '{identifier}' not found in function signature but was requested.")
71
+
72
+ index = self.cached_parameter_indices[identifier]
73
+
74
+ if index >= len(args):
75
+ raise ValueError(f"Expected {index} arguments but got {len(args)}.")
76
+
77
+ return args[index], False, index
78
+
79
+
80
+ def apply_context_parallel(
81
+ module: torch.nn.Module,
82
+ parallel_config: ContextParallelConfig,
83
+ plan: dict[str, ContextParallelModelPlan],
84
+ ) -> None:
85
+ """Apply context parallel on a model."""
86
+ logger.debug(f"Applying context parallel with CP mesh: {parallel_config._mesh} and plan: {plan}")
87
+
88
+ for module_id, cp_model_plan in plan.items():
89
+ submodule = _get_submodule_by_name(module, module_id)
90
+ if not isinstance(submodule, list):
91
+ submodule = [submodule]
92
+
93
+ logger.debug(f"Applying ContextParallelHook to {module_id=} identifying a total of {len(submodule)} modules")
94
+
95
+ for m in submodule:
96
+ if isinstance(cp_model_plan, dict):
97
+ hook = ContextParallelSplitHook(cp_model_plan, parallel_config)
98
+ hook_name = _CONTEXT_PARALLEL_INPUT_HOOK_TEMPLATE.format(module_id)
99
+ elif isinstance(cp_model_plan, (ContextParallelOutput, list, tuple)):
100
+ if isinstance(cp_model_plan, ContextParallelOutput):
101
+ cp_model_plan = [cp_model_plan]
102
+ if not all(isinstance(x, ContextParallelOutput) for x in cp_model_plan):
103
+ raise ValueError(f"Expected all elements of cp_model_plan to be CPOutput, but got {cp_model_plan}")
104
+ hook = ContextParallelGatherHook(cp_model_plan, parallel_config)
105
+ hook_name = _CONTEXT_PARALLEL_OUTPUT_HOOK_TEMPLATE.format(module_id)
106
+ else:
107
+ raise ValueError(f"Unsupported context parallel model plan type: {type(cp_model_plan)}")
108
+ registry = HookRegistry.check_if_exists_or_initialize(m)
109
+ registry.register_hook(hook, hook_name)
110
+
111
+
112
+ def remove_context_parallel(module: torch.nn.Module, plan: dict[str, ContextParallelModelPlan]) -> None:
113
+ for module_id, cp_model_plan in plan.items():
114
+ submodule = _get_submodule_by_name(module, module_id)
115
+ if not isinstance(submodule, list):
116
+ submodule = [submodule]
117
+
118
+ for m in submodule:
119
+ registry = HookRegistry.check_if_exists_or_initialize(m)
120
+ if isinstance(cp_model_plan, dict):
121
+ hook_name = _CONTEXT_PARALLEL_INPUT_HOOK_TEMPLATE.format(module_id)
122
+ elif isinstance(cp_model_plan, (ContextParallelOutput, list, tuple)):
123
+ hook_name = _CONTEXT_PARALLEL_OUTPUT_HOOK_TEMPLATE.format(module_id)
124
+ else:
125
+ raise ValueError(f"Unsupported context parallel model plan type: {type(cp_model_plan)}")
126
+ registry.remove_hook(hook_name)
127
+
128
+
129
+ class ContextParallelSplitHook(ModelHook):
130
+ def __init__(self, metadata: ContextParallelModelPlan, parallel_config: ContextParallelConfig) -> None:
131
+ super().__init__()
132
+ self.metadata = metadata
133
+ self.parallel_config = parallel_config
134
+ self.module_forward_metadata = None
135
+
136
+ def initialize_hook(self, module):
137
+ cls = unwrap_module(module).__class__
138
+ self.module_forward_metadata = ModuleForwardMetadata(_cls=cls)
139
+ return module
140
+
141
+ def pre_forward(self, module, *args, **kwargs):
142
+ args_list = list(args)
143
+
144
+ for name, cpm in self.metadata.items():
145
+ if isinstance(cpm, ContextParallelInput) and cpm.split_output:
146
+ continue
147
+
148
+ # Maybe the parameter was passed as a keyword argument
149
+ input_val, is_kwarg, index = self.module_forward_metadata._get_parameter_from_args_kwargs(
150
+ name, args_list, kwargs
151
+ )
152
+
153
+ if input_val is None:
154
+ continue
155
+
156
+ # The input_val may be a tensor or list/tuple of tensors. In certain cases, user may specify to shard
157
+ # the output instead of input for a particular layer by setting split_output=True
158
+ if isinstance(input_val, torch.Tensor):
159
+ input_val = self._prepare_cp_input(input_val, cpm)
160
+ elif isinstance(input_val, (list, tuple)):
161
+ if len(input_val) != len(cpm):
162
+ raise ValueError(
163
+ f"Expected input model plan to have {len(input_val)} elements, but got {len(cpm)}."
164
+ )
165
+ sharded_input_val = []
166
+ for i, x in enumerate(input_val):
167
+ if torch.is_tensor(x) and not cpm[i].split_output:
168
+ x = self._prepare_cp_input(x, cpm[i])
169
+ sharded_input_val.append(x)
170
+ input_val = sharded_input_val
171
+ else:
172
+ raise ValueError(f"Unsupported input type: {type(input_val)}")
173
+
174
+ if is_kwarg:
175
+ kwargs[name] = input_val
176
+ elif index is not None and index < len(args_list):
177
+ args_list[index] = input_val
178
+ else:
179
+ raise ValueError(
180
+ f"An unexpected error occurred while processing the input '{name}'. Please open an "
181
+ f"issue at https://github.com/huggingface/diffusers/issues and provide a minimal reproducible "
182
+ f"example along with the full stack trace."
183
+ )
184
+
185
+ return tuple(args_list), kwargs
186
+
187
+ def post_forward(self, module, output):
188
+ is_tensor = isinstance(output, torch.Tensor)
189
+ is_tensor_list = isinstance(output, (list, tuple)) and all(isinstance(x, torch.Tensor) for x in output)
190
+
191
+ if not is_tensor and not is_tensor_list:
192
+ raise ValueError(f"Expected output to be a tensor or a list/tuple of tensors, but got {type(output)}.")
193
+
194
+ output = [output] if is_tensor else list(output)
195
+ for index, cpm in self.metadata.items():
196
+ if not isinstance(cpm, ContextParallelInput) or not cpm.split_output:
197
+ continue
198
+ if index >= len(output):
199
+ raise ValueError(f"Index {index} out of bounds for output of length {len(output)}.")
200
+ current_output = output[index]
201
+ current_output = self._prepare_cp_input(current_output, cpm)
202
+ output[index] = current_output
203
+
204
+ return output[0] if is_tensor else tuple(output)
205
+
206
+ def _prepare_cp_input(self, x: torch.Tensor, cp_input: ContextParallelInput) -> torch.Tensor:
207
+ if cp_input.expected_dims is not None and x.dim() != cp_input.expected_dims:
208
+ logger.warning_once(
209
+ f"Expected input tensor to have {cp_input.expected_dims} dimensions, but got {x.dim()} dimensions, split will not be applied."
210
+ )
211
+ return x
212
+ else:
213
+ if self.parallel_config.ulysses_anything or self.parallel_config.ring_anything:
214
+ return PartitionAnythingSharder.shard_anything(
215
+ x, cp_input.split_dim, self.parallel_config._flattened_mesh
216
+ )
217
+ return EquipartitionSharder.shard(x, cp_input.split_dim, self.parallel_config._flattened_mesh)
218
+
219
+
220
+ class ContextParallelGatherHook(ModelHook):
221
+ def __init__(self, metadata: ContextParallelModelPlan, parallel_config: ContextParallelConfig) -> None:
222
+ super().__init__()
223
+ self.metadata = metadata
224
+ self.parallel_config = parallel_config
225
+
226
+ def post_forward(self, module, output):
227
+ is_tensor = isinstance(output, torch.Tensor)
228
+
229
+ if is_tensor:
230
+ output = [output]
231
+ elif not (isinstance(output, (list, tuple)) and all(isinstance(x, torch.Tensor) for x in output)):
232
+ raise ValueError(f"Expected output to be a tensor or a list/tuple of tensors, but got {type(output)}.")
233
+
234
+ output = list(output)
235
+
236
+ if len(output) != len(self.metadata):
237
+ raise ValueError(f"Expected output to have {len(self.metadata)} elements, but got {len(output)}.")
238
+
239
+ for i, cpm in enumerate(self.metadata):
240
+ if cpm is None:
241
+ continue
242
+ if self.parallel_config.ulysses_anything or self.parallel_config.ring_anything:
243
+ output[i] = PartitionAnythingSharder.unshard_anything(
244
+ output[i], cpm.gather_dim, self.parallel_config._flattened_mesh
245
+ )
246
+ else:
247
+ output[i] = EquipartitionSharder.unshard(
248
+ output[i], cpm.gather_dim, self.parallel_config._flattened_mesh
249
+ )
250
+
251
+ return output[0] if is_tensor else tuple(output)
252
+
253
+
254
+ class AllGatherFunction(torch.autograd.Function):
255
+ @staticmethod
256
+ def forward(ctx, tensor, dim, group):
257
+ ctx.dim = dim
258
+ ctx.group = group
259
+ ctx.world_size = torch.distributed.get_world_size(group)
260
+ ctx.rank = torch.distributed.get_rank(group)
261
+ return funcol.all_gather_tensor(tensor, dim, group=group)
262
+
263
+ @staticmethod
264
+ def backward(ctx, grad_output):
265
+ grad_chunks = torch.chunk(grad_output, ctx.world_size, dim=ctx.dim)
266
+ return grad_chunks[ctx.rank], None, None
267
+
268
+
269
+ class EquipartitionSharder:
270
+ @classmethod
271
+ def shard(cls, tensor: torch.Tensor, dim: int, mesh: torch.distributed.device_mesh.DeviceMesh) -> torch.Tensor:
272
+ # NOTE: the following assertion does not have to be true in general. We simply enforce it for now
273
+ # because the alternate case has not yet been tested/required for any model.
274
+ assert tensor.size()[dim] % mesh.size() == 0, (
275
+ "Tensor size along dimension to be sharded must be divisible by mesh size"
276
+ )
277
+
278
+ # The following is not fullgraph compatible with Dynamo (fails in DeviceMesh.get_rank)
279
+ # return tensor.chunk(mesh.size(), dim=dim)[mesh.get_rank()]
280
+
281
+ return tensor.chunk(mesh.size(), dim=dim)[torch.distributed.get_rank(mesh.get_group())]
282
+
283
+ @classmethod
284
+ def unshard(cls, tensor: torch.Tensor, dim: int, mesh: torch.distributed.device_mesh.DeviceMesh) -> torch.Tensor:
285
+ tensor = tensor.contiguous()
286
+ tensor = AllGatherFunction.apply(tensor, dim, mesh.get_group())
287
+ return tensor
288
+
289
+
290
+ class AllGatherAnythingFunction(torch.autograd.Function):
291
+ @staticmethod
292
+ def forward(ctx, tensor: torch.Tensor, dim: int, group: dist.device_mesh.DeviceMesh):
293
+ ctx.dim = dim
294
+ ctx.group = group
295
+ ctx.world_size = dist.get_world_size(group)
296
+ ctx.rank = dist.get_rank(group)
297
+ gathered_tensor = _all_gather_anything(tensor, dim, group)
298
+ return gathered_tensor
299
+
300
+ @staticmethod
301
+ def backward(ctx, grad_output):
302
+ # NOTE: We use `tensor_split` instead of chunk, because the `chunk`
303
+ # function may return fewer than the specified number of chunks!
304
+ grad_splits = torch.tensor_split(grad_output, ctx.world_size, dim=ctx.dim)
305
+ return grad_splits[ctx.rank], None, None
306
+
307
+
308
+ class PartitionAnythingSharder:
309
+ @classmethod
310
+ def shard_anything(
311
+ cls, tensor: torch.Tensor, dim: int, mesh: torch.distributed.device_mesh.DeviceMesh
312
+ ) -> torch.Tensor:
313
+ assert tensor.size()[dim] >= mesh.size(), (
314
+ f"Cannot shard tensor of size {tensor.size()} along dim {dim} across mesh of size {mesh.size()}."
315
+ )
316
+ # NOTE: We use `tensor_split` instead of chunk, because the `chunk`
317
+ # function may return fewer than the specified number of chunks!
318
+ return tensor.tensor_split(mesh.size(), dim=dim)[dist.get_rank(mesh.get_group())]
319
+
320
+ @classmethod
321
+ def unshard_anything(
322
+ cls, tensor: torch.Tensor, dim: int, mesh: torch.distributed.device_mesh.DeviceMesh
323
+ ) -> torch.Tensor:
324
+ tensor = tensor.contiguous()
325
+ tensor = AllGatherAnythingFunction.apply(tensor, dim, mesh.get_group())
326
+ return tensor
327
+
328
+
329
+ @lru_cache_unless_export(maxsize=64)
330
+ def _fill_gather_shapes(shape: tuple[int], gather_dims: tuple[int], dim: int, world_size: int) -> list[list[int]]:
331
+ gather_shapes = []
332
+ for i in range(world_size):
333
+ rank_shape = list(copy.deepcopy(shape))
334
+ rank_shape[dim] = gather_dims[i]
335
+ gather_shapes.append(rank_shape)
336
+ return gather_shapes
337
+
338
+
339
+ @maybe_allow_in_graph
340
+ def _all_gather_anything(tensor: torch.Tensor, dim: int, group: dist.device_mesh.DeviceMesh) -> torch.Tensor:
341
+ world_size = dist.get_world_size(group=group)
342
+
343
+ tensor = tensor.contiguous()
344
+ shape = tensor.shape
345
+ rank_dim = shape[dim]
346
+ gather_dims = gather_size_by_comm(rank_dim, group)
347
+
348
+ gather_shapes = _fill_gather_shapes(tuple(shape), tuple(gather_dims), dim, world_size)
349
+
350
+ gathered_tensors = [torch.empty(shape, device=tensor.device, dtype=tensor.dtype) for shape in gather_shapes]
351
+
352
+ dist.all_gather(gathered_tensors, tensor, group=group)
353
+ gathered_tensor = torch.cat(gathered_tensors, dim=dim)
354
+ return gathered_tensor
355
+
356
+
357
+ def _get_submodule_by_name(model: torch.nn.Module, name: str) -> torch.nn.Module | list[torch.nn.Module]:
358
+ if name.count("*") > 1:
359
+ raise ValueError("Wildcard '*' can only be used once in the name")
360
+ return _find_submodule_by_name(model, name)
361
+
362
+
363
+ def _find_submodule_by_name(model: torch.nn.Module, name: str) -> torch.nn.Module | list[torch.nn.Module]:
364
+ if name == "":
365
+ return model
366
+ first_atom, remaining_name = name.split(".", 1) if "." in name else (name, "")
367
+ if first_atom == "*":
368
+ if not isinstance(model, torch.nn.ModuleList):
369
+ raise ValueError("Wildcard '*' can only be used with ModuleList")
370
+ submodules = []
371
+ for submodule in model:
372
+ subsubmodules = _find_submodule_by_name(submodule, remaining_name)
373
+ if not isinstance(subsubmodules, list):
374
+ subsubmodules = [subsubmodules]
375
+ submodules.extend(subsubmodules)
376
+ return submodules
377
+ else:
378
+ if hasattr(model, first_atom):
379
+ submodule = getattr(model, first_atom)
380
+ return _find_submodule_by_name(submodule, remaining_name)
381
+ else:
382
+ raise ValueError(f"'{first_atom}' is not a submodule of '{model.__class__.__name__}'")
diffusers/hooks/faster_cache.py ADDED
@@ -0,0 +1,654 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 The HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ import re
16
+ from dataclasses import dataclass
17
+ from typing import Any, Callable
18
+
19
+ import torch
20
+
21
+ from ..models.attention import AttentionModuleMixin
22
+ from ..models.modeling_outputs import Transformer2DModelOutput
23
+ from ..utils import logging
24
+ from ._common import _ATTENTION_CLASSES
25
+ from .hooks import HookRegistry, ModelHook
26
+
27
+
28
+ logger = logging.get_logger(__name__) # pylint: disable=invalid-name
29
+
30
+
31
+ _FASTER_CACHE_DENOISER_HOOK = "faster_cache_denoiser"
32
+ _FASTER_CACHE_BLOCK_HOOK = "faster_cache_block"
33
+ _SPATIAL_ATTENTION_BLOCK_IDENTIFIERS = (
34
+ "^blocks.*attn",
35
+ "^transformer_blocks.*attn",
36
+ "^single_transformer_blocks.*attn",
37
+ )
38
+ _TEMPORAL_ATTENTION_BLOCK_IDENTIFIERS = ("^temporal_transformer_blocks.*attn",)
39
+ _TRANSFORMER_BLOCK_IDENTIFIERS = _SPATIAL_ATTENTION_BLOCK_IDENTIFIERS + _TEMPORAL_ATTENTION_BLOCK_IDENTIFIERS
40
+ _UNCOND_COND_INPUT_KWARGS_IDENTIFIERS = (
41
+ "hidden_states",
42
+ "encoder_hidden_states",
43
+ "timestep",
44
+ "attention_mask",
45
+ "encoder_attention_mask",
46
+ )
47
+
48
+
49
+ @dataclass
50
+ class FasterCacheConfig:
51
+ r"""
52
+ Configuration for [FasterCache](https://huggingface.co/papers/2410.19355).
53
+
54
+ Attributes:
55
+ spatial_attention_block_skip_range (`int`, defaults to `2`):
56
+ Calculate the attention states every `N` iterations. If this is set to `N`, the attention computation will
57
+ be skipped `N - 1` times (i.e., cached attention states will be reused) before computing the new attention
58
+ states again.
59
+ temporal_attention_block_skip_range (`int`, *optional*, defaults to `None`):
60
+ Calculate the attention states every `N` iterations. If this is set to `N`, the attention computation will
61
+ be skipped `N - 1` times (i.e., cached attention states will be reused) before computing the new attention
62
+ states again.
63
+ spatial_attention_timestep_skip_range (`tuple[float, float]`, defaults to `(-1, 681)`):
64
+ The timestep range within which the spatial attention computation can be skipped without a significant loss
65
+ in quality. This is to be determined by the user based on the underlying model. The first value in the
66
+ tuple is the lower bound and the second value is the upper bound. Typically, diffusion timesteps for
67
+ denoising are in the reversed range of 0 to 1000 (i.e. denoising starts at timestep 1000 and ends at
68
+ timestep 0). For the default values, this would mean that the spatial attention computation skipping will
69
+ be applicable only after denoising timestep 681 is reached, and continue until the end of the denoising
70
+ process.
71
+ temporal_attention_timestep_skip_range (`tuple[float, float]`, *optional*, defaults to `None`):
72
+ The timestep range within which the temporal attention computation can be skipped without a significant
73
+ loss in quality. This is to be determined by the user based on the underlying model. The first value in the
74
+ tuple is the lower bound and the second value is the upper bound. Typically, diffusion timesteps for
75
+ denoising are in the reversed range of 0 to 1000 (i.e. denoising starts at timestep 1000 and ends at
76
+ timestep 0).
77
+ low_frequency_weight_update_timestep_range (`tuple[int, int]`, defaults to `(99, 901)`):
78
+ The timestep range within which the low frequency weight scaling update is applied. The first value in the
79
+ tuple is the lower bound and the second value is the upper bound of the timestep range. The callback
80
+ function for the update is called only within this range.
81
+ high_frequency_weight_update_timestep_range (`tuple[int, int]`, defaults to `(-1, 301)`):
82
+ The timestep range within which the high frequency weight scaling update is applied. The first value in the
83
+ tuple is the lower bound and the second value is the upper bound of the timestep range. The callback
84
+ function for the update is called only within this range.
85
+ alpha_low_frequency (`float`, defaults to `1.1`):
86
+ The weight to scale the low frequency updates by. This is used to approximate the unconditional branch from
87
+ the conditional branch outputs.
88
+ alpha_high_frequency (`float`, defaults to `1.1`):
89
+ The weight to scale the high frequency updates by. This is used to approximate the unconditional branch
90
+ from the conditional branch outputs.
91
+ unconditional_batch_skip_range (`int`, defaults to `5`):
92
+ Process the unconditional branch every `N` iterations. If this is set to `N`, the unconditional branch
93
+ computation will be skipped `N - 1` times (i.e., cached unconditional branch states will be reused) before
94
+ computing the new unconditional branch states again.
95
+ unconditional_batch_timestep_skip_range (`tuple[float, float]`, defaults to `(-1, 641)`):
96
+ The timestep range within which the unconditional branch computation can be skipped without a significant
97
+ loss in quality. This is to be determined by the user based on the underlying model. The first value in the
98
+ tuple is the lower bound and the second value is the upper bound.
99
+ spatial_attention_block_identifiers (`tuple[str, ...]`, defaults to `("blocks.*attn1", "transformer_blocks.*attn1", "single_transformer_blocks.*attn1")`):
100
+ The identifiers to match the spatial attention blocks in the model. If the name of the block contains any
101
+ of these identifiers, FasterCache will be applied to that block. This can either be the full layer names,
102
+ partial layer names, or regex patterns. Matching will always be done using a regex match.
103
+ temporal_attention_block_identifiers (`tuple[str, ...]`, defaults to `("temporal_transformer_blocks.*attn1",)`):
104
+ The identifiers to match the temporal attention blocks in the model. If the name of the block contains any
105
+ of these identifiers, FasterCache will be applied to that block. This can either be the full layer names,
106
+ partial layer names, or regex patterns. Matching will always be done using a regex match.
107
+ attention_weight_callback (`Callable[[torch.nn.Module], float]`, defaults to `None`):
108
+ The callback function to determine the weight to scale the attention outputs by. This function should take
109
+ the attention module as input and return a float value. This is used to approximate the unconditional
110
+ branch from the conditional branch outputs. If not provided, the default weight is 0.5 for all timesteps.
111
+ Typically, as described in the paper, this weight should gradually increase from 0 to 1 as the inference
112
+ progresses. Users are encouraged to experiment and provide custom weight schedules that take into account
113
+ the number of inference steps and underlying model behaviour as denoising progresses.
114
+ low_frequency_weight_callback (`Callable[[torch.nn.Module], float]`, defaults to `None`):
115
+ The callback function to determine the weight to scale the low frequency updates by. If not provided, the
116
+ default weight is 1.1 for timesteps within the range specified (as described in the paper).
117
+ high_frequency_weight_callback (`Callable[[torch.nn.Module], float]`, defaults to `None`):
118
+ The callback function to determine the weight to scale the high frequency updates by. If not provided, the
119
+ default weight is 1.1 for timesteps within the range specified (as described in the paper).
120
+ tensor_format (`str`, defaults to `"BCFHW"`):
121
+ The format of the input tensors. This should be one of `"BCFHW"`, `"BFCHW"`, or `"BCHW"`. The format is
122
+ used to split individual latent frames in order for low and high frequency components to be computed.
123
+ is_guidance_distilled (`bool`, defaults to `False`):
124
+ Whether the model is guidance distilled or not. If the model is guidance distilled, FasterCache will not be
125
+ applied at the denoiser-level to skip the unconditional branch computation (as there is none).
126
+ _unconditional_conditional_input_kwargs_identifiers (`list[str]`, defaults to `("hidden_states", "encoder_hidden_states", "timestep", "attention_mask", "encoder_attention_mask")`):
127
+ The identifiers to match the input kwargs that contain the batchwise-concatenated unconditional and
128
+ conditional inputs. If the name of the input kwargs contains any of these identifiers, FasterCache will
129
+ split the inputs into unconditional and conditional branches. This must be a list of exact input kwargs
130
+ names that contain the batchwise-concatenated unconditional and conditional inputs.
131
+ """
132
+
133
+ # In the paper and codebase, they hardcode these values to 2. However, it can be made configurable
134
+ # after some testing. We default to 2 if these parameters are not provided.
135
+ spatial_attention_block_skip_range: int = 2
136
+ temporal_attention_block_skip_range: int | None = None
137
+
138
+ spatial_attention_timestep_skip_range: tuple[int, int] = (-1, 681)
139
+ temporal_attention_timestep_skip_range: tuple[int, int] = (-1, 681)
140
+
141
+ # Indicator functions for low/high frequency as mentioned in Equation 11 of the paper
142
+ low_frequency_weight_update_timestep_range: tuple[int, int] = (99, 901)
143
+ high_frequency_weight_update_timestep_range: tuple[int, int] = (-1, 301)
144
+
145
+ # ⍺1 and ⍺2 as mentioned in Equation 11 of the paper
146
+ alpha_low_frequency: float = 1.1
147
+ alpha_high_frequency: float = 1.1
148
+
149
+ # n as described in CFG-Cache explanation in the paper - dependent on the model
150
+ unconditional_batch_skip_range: int = 5
151
+ unconditional_batch_timestep_skip_range: tuple[int, int] = (-1, 641)
152
+
153
+ spatial_attention_block_identifiers: tuple[str, ...] = _SPATIAL_ATTENTION_BLOCK_IDENTIFIERS
154
+ temporal_attention_block_identifiers: tuple[str, ...] = _TEMPORAL_ATTENTION_BLOCK_IDENTIFIERS
155
+
156
+ attention_weight_callback: Callable[[torch.nn.Module], float] = None
157
+ low_frequency_weight_callback: Callable[[torch.nn.Module], float] = None
158
+ high_frequency_weight_callback: Callable[[torch.nn.Module], float] = None
159
+
160
+ tensor_format: str = "BCFHW"
161
+ is_guidance_distilled: bool = False
162
+
163
+ current_timestep_callback: Callable[[], int] = None
164
+
165
+ _unconditional_conditional_input_kwargs_identifiers: list[str] = _UNCOND_COND_INPUT_KWARGS_IDENTIFIERS
166
+
167
+ def __repr__(self) -> str:
168
+ return (
169
+ f"FasterCacheConfig(\n"
170
+ f" spatial_attention_block_skip_range={self.spatial_attention_block_skip_range},\n"
171
+ f" temporal_attention_block_skip_range={self.temporal_attention_block_skip_range},\n"
172
+ f" spatial_attention_timestep_skip_range={self.spatial_attention_timestep_skip_range},\n"
173
+ f" temporal_attention_timestep_skip_range={self.temporal_attention_timestep_skip_range},\n"
174
+ f" low_frequency_weight_update_timestep_range={self.low_frequency_weight_update_timestep_range},\n"
175
+ f" high_frequency_weight_update_timestep_range={self.high_frequency_weight_update_timestep_range},\n"
176
+ f" alpha_low_frequency={self.alpha_low_frequency},\n"
177
+ f" alpha_high_frequency={self.alpha_high_frequency},\n"
178
+ f" unconditional_batch_skip_range={self.unconditional_batch_skip_range},\n"
179
+ f" unconditional_batch_timestep_skip_range={self.unconditional_batch_timestep_skip_range},\n"
180
+ f" spatial_attention_block_identifiers={self.spatial_attention_block_identifiers},\n"
181
+ f" temporal_attention_block_identifiers={self.temporal_attention_block_identifiers},\n"
182
+ f" tensor_format={self.tensor_format},\n"
183
+ f")"
184
+ )
185
+
186
+
187
+ class FasterCacheDenoiserState:
188
+ r"""
189
+ State for [FasterCache](https://huggingface.co/papers/2410.19355) top-level denoiser module.
190
+ """
191
+
192
+ def __init__(self) -> None:
193
+ self.iteration: int = 0
194
+ self.low_frequency_delta: torch.Tensor = None
195
+ self.high_frequency_delta: torch.Tensor = None
196
+
197
+ def reset(self):
198
+ self.iteration = 0
199
+ self.low_frequency_delta = None
200
+ self.high_frequency_delta = None
201
+
202
+
203
+ class FasterCacheBlockState:
204
+ r"""
205
+ State for [FasterCache](https://huggingface.co/papers/2410.19355). Every underlying block that FasterCache is
206
+ applied to will have an instance of this state.
207
+ """
208
+
209
+ def __init__(self) -> None:
210
+ self.iteration: int = 0
211
+ self.batch_size: int = None
212
+ self.cache: tuple[torch.Tensor, torch.Tensor] = None
213
+
214
+ def reset(self):
215
+ self.iteration = 0
216
+ self.batch_size = None
217
+ self.cache = None
218
+
219
+
220
+ class FasterCacheDenoiserHook(ModelHook):
221
+ _is_stateful = True
222
+
223
+ def __init__(
224
+ self,
225
+ unconditional_batch_skip_range: int,
226
+ unconditional_batch_timestep_skip_range: tuple[int, int],
227
+ tensor_format: str,
228
+ is_guidance_distilled: bool,
229
+ uncond_cond_input_kwargs_identifiers: list[str],
230
+ current_timestep_callback: Callable[[], int],
231
+ low_frequency_weight_callback: Callable[[torch.nn.Module], torch.Tensor],
232
+ high_frequency_weight_callback: Callable[[torch.nn.Module], torch.Tensor],
233
+ ) -> None:
234
+ super().__init__()
235
+
236
+ self.unconditional_batch_skip_range = unconditional_batch_skip_range
237
+ self.unconditional_batch_timestep_skip_range = unconditional_batch_timestep_skip_range
238
+ # We can't easily detect what args are to be split in unconditional and conditional branches. We
239
+ # can only do it for kwargs, hence they are the only ones we split. The args are passed as-is.
240
+ # If a model is to be made compatible with FasterCache, the user must ensure that the inputs that
241
+ # contain batchwise-concatenated unconditional and conditional inputs are passed as kwargs.
242
+ self.uncond_cond_input_kwargs_identifiers = uncond_cond_input_kwargs_identifiers
243
+ self.tensor_format = tensor_format
244
+ self.is_guidance_distilled = is_guidance_distilled
245
+
246
+ self.current_timestep_callback = current_timestep_callback
247
+ self.low_frequency_weight_callback = low_frequency_weight_callback
248
+ self.high_frequency_weight_callback = high_frequency_weight_callback
249
+
250
+ def initialize_hook(self, module):
251
+ self.state = FasterCacheDenoiserState()
252
+ return module
253
+
254
+ @staticmethod
255
+ def _get_cond_input(input: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
256
+ # Note: this method assumes that the input tensor is batchwise-concatenated with unconditional inputs
257
+ # followed by conditional inputs.
258
+ _, cond = input.chunk(2, dim=0)
259
+ return cond
260
+
261
+ def new_forward(self, module: torch.nn.Module, *args, **kwargs) -> Any:
262
+ # Split the unconditional and conditional inputs. We only want to infer the conditional branch if the
263
+ # requirements for skipping the unconditional branch are met as described in the paper.
264
+ # We skip the unconditional branch only if the following conditions are met:
265
+ # 1. We have completed at least one iteration of the denoiser
266
+ # 2. The current timestep is within the range specified by the user. This is the optimal timestep range
267
+ # where approximating the unconditional branch from the computation of the conditional branch is possible
268
+ # without a significant loss in quality.
269
+ # 3. The current iteration is not a multiple of the unconditional batch skip range. This is done so that
270
+ # we compute the unconditional branch at least once every few iterations to ensure minimal quality loss.
271
+ is_within_timestep_range = (
272
+ self.unconditional_batch_timestep_skip_range[0]
273
+ < self.current_timestep_callback()
274
+ < self.unconditional_batch_timestep_skip_range[1]
275
+ )
276
+ should_skip_uncond = (
277
+ self.state.iteration > 0
278
+ and is_within_timestep_range
279
+ and self.state.iteration % self.unconditional_batch_skip_range != 0
280
+ and not self.is_guidance_distilled
281
+ )
282
+
283
+ if should_skip_uncond:
284
+ is_any_kwarg_uncond = any(k in self.uncond_cond_input_kwargs_identifiers for k in kwargs.keys())
285
+ if is_any_kwarg_uncond:
286
+ logger.debug("FasterCache - Skipping unconditional branch computation")
287
+ args = tuple([self._get_cond_input(arg) if torch.is_tensor(arg) else arg for arg in args])
288
+ kwargs = {
289
+ k: v if k not in self.uncond_cond_input_kwargs_identifiers else self._get_cond_input(v)
290
+ for k, v in kwargs.items()
291
+ }
292
+
293
+ output = self.fn_ref.original_forward(*args, **kwargs)
294
+
295
+ if self.is_guidance_distilled:
296
+ self.state.iteration += 1
297
+ return output
298
+
299
+ if torch.is_tensor(output):
300
+ hidden_states = output
301
+ elif isinstance(output, (tuple, Transformer2DModelOutput)):
302
+ hidden_states = output[0]
303
+
304
+ batch_size = hidden_states.size(0)
305
+
306
+ if should_skip_uncond:
307
+ self.state.low_frequency_delta = self.state.low_frequency_delta * self.low_frequency_weight_callback(
308
+ module
309
+ )
310
+ self.state.high_frequency_delta = self.state.high_frequency_delta * self.high_frequency_weight_callback(
311
+ module
312
+ )
313
+
314
+ if self.tensor_format == "BCFHW":
315
+ hidden_states = hidden_states.permute(0, 2, 1, 3, 4)
316
+ if self.tensor_format == "BCFHW" or self.tensor_format == "BFCHW":
317
+ hidden_states = hidden_states.flatten(0, 1)
318
+
319
+ low_freq_cond, high_freq_cond = _split_low_high_freq(hidden_states.float())
320
+
321
+ # Approximate/compute the unconditional branch outputs as described in Equation 9 and 10 of the paper
322
+ low_freq_uncond = self.state.low_frequency_delta + low_freq_cond
323
+ high_freq_uncond = self.state.high_frequency_delta + high_freq_cond
324
+ uncond_freq = low_freq_uncond + high_freq_uncond
325
+
326
+ uncond_states = torch.fft.ifftshift(uncond_freq)
327
+ uncond_states = torch.fft.ifft2(uncond_states).real
328
+
329
+ if self.tensor_format == "BCFHW" or self.tensor_format == "BFCHW":
330
+ uncond_states = uncond_states.unflatten(0, (batch_size, -1))
331
+ hidden_states = hidden_states.unflatten(0, (batch_size, -1))
332
+ if self.tensor_format == "BCFHW":
333
+ uncond_states = uncond_states.permute(0, 2, 1, 3, 4)
334
+ hidden_states = hidden_states.permute(0, 2, 1, 3, 4)
335
+
336
+ # Concatenate the approximated unconditional and predicted conditional branches
337
+ uncond_states = uncond_states.to(hidden_states.dtype)
338
+ hidden_states = torch.cat([uncond_states, hidden_states], dim=0)
339
+ else:
340
+ uncond_states, cond_states = hidden_states.chunk(2, dim=0)
341
+ if self.tensor_format == "BCFHW":
342
+ uncond_states = uncond_states.permute(0, 2, 1, 3, 4)
343
+ cond_states = cond_states.permute(0, 2, 1, 3, 4)
344
+ if self.tensor_format == "BCFHW" or self.tensor_format == "BFCHW":
345
+ uncond_states = uncond_states.flatten(0, 1)
346
+ cond_states = cond_states.flatten(0, 1)
347
+
348
+ low_freq_uncond, high_freq_uncond = _split_low_high_freq(uncond_states.float())
349
+ low_freq_cond, high_freq_cond = _split_low_high_freq(cond_states.float())
350
+ self.state.low_frequency_delta = low_freq_uncond - low_freq_cond
351
+ self.state.high_frequency_delta = high_freq_uncond - high_freq_cond
352
+
353
+ self.state.iteration += 1
354
+ if torch.is_tensor(output):
355
+ output = hidden_states
356
+ elif isinstance(output, tuple):
357
+ output = (hidden_states, *output[1:])
358
+ else:
359
+ output.sample = hidden_states
360
+
361
+ return output
362
+
363
+ def reset_state(self, module: torch.nn.Module) -> torch.nn.Module:
364
+ self.state.reset()
365
+ return module
366
+
367
+
368
+ class FasterCacheBlockHook(ModelHook):
369
+ _is_stateful = True
370
+
371
+ def __init__(
372
+ self,
373
+ block_skip_range: int,
374
+ timestep_skip_range: tuple[int, int],
375
+ is_guidance_distilled: bool,
376
+ weight_callback: Callable[[torch.nn.Module], float],
377
+ current_timestep_callback: Callable[[], int],
378
+ ) -> None:
379
+ super().__init__()
380
+
381
+ self.block_skip_range = block_skip_range
382
+ self.timestep_skip_range = timestep_skip_range
383
+ self.is_guidance_distilled = is_guidance_distilled
384
+
385
+ self.weight_callback = weight_callback
386
+ self.current_timestep_callback = current_timestep_callback
387
+
388
+ def initialize_hook(self, module):
389
+ self.state = FasterCacheBlockState()
390
+ return module
391
+
392
+ def _compute_approximated_attention_output(
393
+ self, t_2_output: torch.Tensor, t_output: torch.Tensor, weight: float, batch_size: int
394
+ ) -> torch.Tensor:
395
+ if t_2_output.size(0) != batch_size:
396
+ # The cache t_2_output contains both batchwise-concatenated unconditional-conditional branch outputs. Just
397
+ # take the conditional branch outputs.
398
+ assert t_2_output.size(0) == 2 * batch_size
399
+ t_2_output = t_2_output[batch_size:]
400
+ if t_output.size(0) != batch_size:
401
+ # The cache t_output contains both batchwise-concatenated unconditional-conditional branch outputs. Just
402
+ # take the conditional branch outputs.
403
+ assert t_output.size(0) == 2 * batch_size
404
+ t_output = t_output[batch_size:]
405
+ return t_output + (t_output - t_2_output) * weight
406
+
407
+ def new_forward(self, module: torch.nn.Module, *args, **kwargs) -> Any:
408
+ batch_size = [
409
+ *[arg.size(0) for arg in args if torch.is_tensor(arg)],
410
+ *[v.size(0) for v in kwargs.values() if torch.is_tensor(v)],
411
+ ][0]
412
+ if self.state.batch_size is None:
413
+ # Will be updated on first forward pass through the denoiser
414
+ self.state.batch_size = batch_size
415
+
416
+ # If we have to skip due to the skip conditions, then let's skip as expected.
417
+ # But, we can't skip if the denoiser wants to infer both unconditional and conditional branches. This
418
+ # is because the expected output shapes of attention layer will not match if we only return values from
419
+ # the cache (which only caches conditional branch outputs). So, if state.batch_size (which is the true
420
+ # unconditional-conditional batch size) is same as the current batch size, we don't perform the layer
421
+ # skip. Otherwise, we conditionally skip the layer based on what state.skip_callback returns.
422
+ is_within_timestep_range = (
423
+ self.timestep_skip_range[0] < self.current_timestep_callback() < self.timestep_skip_range[1]
424
+ )
425
+ if not is_within_timestep_range:
426
+ should_skip_attention = False
427
+ else:
428
+ should_compute_attention = self.state.iteration > 0 and self.state.iteration % self.block_skip_range == 0
429
+ should_skip_attention = not should_compute_attention
430
+ if should_skip_attention:
431
+ should_skip_attention = self.is_guidance_distilled or self.state.batch_size != batch_size
432
+
433
+ if should_skip_attention:
434
+ logger.debug("FasterCache - Skipping attention and using approximation")
435
+ if torch.is_tensor(self.state.cache[-1]):
436
+ t_2_output, t_output = self.state.cache
437
+ weight = self.weight_callback(module)
438
+ output = self._compute_approximated_attention_output(t_2_output, t_output, weight, batch_size)
439
+ else:
440
+ # The cache contains multiple tensors from past N iterations (N=2 for FasterCache). We need to handle all of them.
441
+ # Diffusers blocks can return multiple tensors - let's call them [A, B, C, ...] for simplicity.
442
+ # In our cache, we would have [[A_1, B_1, C_1, ...], [A_2, B_2, C_2, ...], ...] where each list is the output from
443
+ # a forward pass of the block. We need to compute the approximated output for each of these tensors.
444
+ # The zip(*state.cache) operation will give us [(A_1, A_2, ...), (B_1, B_2, ...), (C_1, C_2, ...), ...] which
445
+ # allows us to compute the approximated attention output for each tensor in the cache.
446
+ output = ()
447
+ for t_2_output, t_output in zip(*self.state.cache):
448
+ result = self._compute_approximated_attention_output(
449
+ t_2_output, t_output, self.weight_callback(module), batch_size
450
+ )
451
+ output += (result,)
452
+ else:
453
+ logger.debug("FasterCache - Computing attention")
454
+ output = self.fn_ref.original_forward(*args, **kwargs)
455
+
456
+ # Note that the following condition for getting hidden_states should suffice since Diffusers blocks either return
457
+ # a single hidden_states tensor, or a tuple of (hidden_states, encoder_hidden_states) tensors. We need to handle
458
+ # both cases.
459
+ if torch.is_tensor(output):
460
+ cache_output = output
461
+ if not self.is_guidance_distilled and cache_output.size(0) == self.state.batch_size:
462
+ # The output here can be both unconditional-conditional branch outputs or just conditional branch outputs.
463
+ # This is determined at the higher-level denoiser module. We only want to cache the conditional branch outputs.
464
+ cache_output = cache_output.chunk(2, dim=0)[1]
465
+ else:
466
+ # Cache all return values and perform the same operation as above
467
+ cache_output = ()
468
+ for out in output:
469
+ if not self.is_guidance_distilled and out.size(0) == self.state.batch_size:
470
+ out = out.chunk(2, dim=0)[1]
471
+ cache_output += (out,)
472
+
473
+ if self.state.cache is None:
474
+ self.state.cache = [cache_output, cache_output]
475
+ else:
476
+ self.state.cache = [self.state.cache[-1], cache_output]
477
+
478
+ self.state.iteration += 1
479
+ return output
480
+
481
+ def reset_state(self, module: torch.nn.Module) -> torch.nn.Module:
482
+ self.state.reset()
483
+ return module
484
+
485
+
486
+ def apply_faster_cache(module: torch.nn.Module, config: FasterCacheConfig) -> None:
487
+ r"""
488
+ Applies [FasterCache](https://huggingface.co/papers/2410.19355) to a given pipeline.
489
+
490
+ Args:
491
+ module (`torch.nn.Module`):
492
+ The pytorch module to apply FasterCache to. Typically, this should be a transformer architecture supported
493
+ in Diffusers, such as `CogVideoXTransformer3DModel`, but external implementations may also work.
494
+ config (`FasterCacheConfig`):
495
+ The configuration to use for FasterCache.
496
+
497
+ Example:
498
+ ```python
499
+ >>> import torch
500
+ >>> from diffusers import CogVideoXPipeline, FasterCacheConfig, apply_faster_cache
501
+
502
+ >>> pipe = CogVideoXPipeline.from_pretrained("THUDM/CogVideoX-5b", torch_dtype=torch.bfloat16)
503
+ >>> pipe.to("cuda")
504
+
505
+ >>> config = FasterCacheConfig(
506
+ ... spatial_attention_block_skip_range=2,
507
+ ... spatial_attention_timestep_skip_range=(-1, 681),
508
+ ... low_frequency_weight_update_timestep_range=(99, 641),
509
+ ... high_frequency_weight_update_timestep_range=(-1, 301),
510
+ ... spatial_attention_block_identifiers=["transformer_blocks"],
511
+ ... attention_weight_callback=lambda _: 0.3,
512
+ ... tensor_format="BFCHW",
513
+ ... )
514
+ >>> apply_faster_cache(pipe.transformer, config)
515
+ ```
516
+ """
517
+
518
+ logger.warning(
519
+ "FasterCache is a purely experimental feature and may not work as expected. Not all models support FasterCache. "
520
+ "The API is subject to change in future releases, with no guarantee of backward compatibility. Please report any issues at "
521
+ "https://github.com/huggingface/diffusers/issues."
522
+ )
523
+
524
+ if config.attention_weight_callback is None:
525
+ # If the user has not provided a weight callback, we default to 0.5 for all timesteps.
526
+ # In the paper, they recommend using a gradually increasing weight from 0 to 1 as the inference progresses, but
527
+ # this depends from model-to-model. It is required by the user to provide a weight callback if they want to
528
+ # use a different weight function. Defaulting to 0.5 works well in practice for most cases.
529
+ logger.warning(
530
+ "No `attention_weight_callback` provided when enabling FasterCache. Defaulting to using a weight of 0.5 for all timesteps."
531
+ )
532
+ config.attention_weight_callback = lambda _: 0.5
533
+
534
+ if config.low_frequency_weight_callback is None:
535
+ logger.debug(
536
+ "Low frequency weight callback not provided when enabling FasterCache. Defaulting to behaviour described in the paper."
537
+ )
538
+
539
+ def low_frequency_weight_callback(module: torch.nn.Module) -> float:
540
+ is_within_range = (
541
+ config.low_frequency_weight_update_timestep_range[0]
542
+ < config.current_timestep_callback()
543
+ < config.low_frequency_weight_update_timestep_range[1]
544
+ )
545
+ return config.alpha_low_frequency if is_within_range else 1.0
546
+
547
+ config.low_frequency_weight_callback = low_frequency_weight_callback
548
+
549
+ if config.high_frequency_weight_callback is None:
550
+ logger.debug(
551
+ "High frequency weight callback not provided when enabling FasterCache. Defaulting to behaviour described in the paper."
552
+ )
553
+
554
+ def high_frequency_weight_callback(module: torch.nn.Module) -> float:
555
+ is_within_range = (
556
+ config.high_frequency_weight_update_timestep_range[0]
557
+ < config.current_timestep_callback()
558
+ < config.high_frequency_weight_update_timestep_range[1]
559
+ )
560
+ return config.alpha_high_frequency if is_within_range else 1.0
561
+
562
+ config.high_frequency_weight_callback = high_frequency_weight_callback
563
+
564
+ supported_tensor_formats = ["BCFHW", "BFCHW", "BCHW"] # TODO(aryan): Support BSC for LTX Video
565
+ if config.tensor_format not in supported_tensor_formats:
566
+ raise ValueError(f"`tensor_format` must be one of {supported_tensor_formats}, but got {config.tensor_format}.")
567
+
568
+ _apply_faster_cache_on_denoiser(module, config)
569
+
570
+ for name, submodule in module.named_modules():
571
+ if not isinstance(submodule, _ATTENTION_CLASSES):
572
+ continue
573
+ if any(re.search(identifier, name) is not None for identifier in _TRANSFORMER_BLOCK_IDENTIFIERS):
574
+ _apply_faster_cache_on_attention_class(name, submodule, config)
575
+
576
+
577
+ def _apply_faster_cache_on_denoiser(module: torch.nn.Module, config: FasterCacheConfig) -> None:
578
+ hook = FasterCacheDenoiserHook(
579
+ config.unconditional_batch_skip_range,
580
+ config.unconditional_batch_timestep_skip_range,
581
+ config.tensor_format,
582
+ config.is_guidance_distilled,
583
+ config._unconditional_conditional_input_kwargs_identifiers,
584
+ config.current_timestep_callback,
585
+ config.low_frequency_weight_callback,
586
+ config.high_frequency_weight_callback,
587
+ )
588
+ registry = HookRegistry.check_if_exists_or_initialize(module)
589
+ registry.register_hook(hook, _FASTER_CACHE_DENOISER_HOOK)
590
+
591
+
592
+ def _apply_faster_cache_on_attention_class(name: str, module: AttentionModuleMixin, config: FasterCacheConfig) -> None:
593
+ is_spatial_self_attention = (
594
+ any(re.search(identifier, name) is not None for identifier in config.spatial_attention_block_identifiers)
595
+ and config.spatial_attention_block_skip_range is not None
596
+ and not getattr(module, "is_cross_attention", False)
597
+ )
598
+ is_temporal_self_attention = (
599
+ any(re.search(identifier, name) is not None for identifier in config.temporal_attention_block_identifiers)
600
+ and config.temporal_attention_block_skip_range is not None
601
+ and not module.is_cross_attention
602
+ )
603
+
604
+ block_skip_range, timestep_skip_range, block_type = None, None, None
605
+ if is_spatial_self_attention:
606
+ block_skip_range = config.spatial_attention_block_skip_range
607
+ timestep_skip_range = config.spatial_attention_timestep_skip_range
608
+ block_type = "spatial"
609
+ elif is_temporal_self_attention:
610
+ block_skip_range = config.temporal_attention_block_skip_range
611
+ timestep_skip_range = config.temporal_attention_timestep_skip_range
612
+ block_type = "temporal"
613
+
614
+ if block_skip_range is None or timestep_skip_range is None:
615
+ logger.debug(
616
+ f'Unable to apply FasterCache to the selected layer: "{name}" because it does '
617
+ f"not match any of the required criteria for spatial or temporal attention layers. Note, "
618
+ f"however, that this layer may still be valid for applying PAB. Please specify the correct "
619
+ f"block identifiers in the configuration or use the specialized `apply_faster_cache_on_module` "
620
+ f"function to apply FasterCache to this layer."
621
+ )
622
+ return
623
+
624
+ logger.debug(f"Enabling FasterCache ({block_type}) for layer: {name}")
625
+ hook = FasterCacheBlockHook(
626
+ block_skip_range,
627
+ timestep_skip_range,
628
+ config.is_guidance_distilled,
629
+ config.attention_weight_callback,
630
+ config.current_timestep_callback,
631
+ )
632
+ registry = HookRegistry.check_if_exists_or_initialize(module)
633
+ registry.register_hook(hook, _FASTER_CACHE_BLOCK_HOOK)
634
+
635
+
636
+ # Reference: https://github.com/Vchitect/FasterCache/blob/fab32c15014636dc854948319c0a9a8d92c7acb4/scripts/latte/faster_cache_sample_latte.py#L127C1-L143C39
637
+ @torch.no_grad()
638
+ def _split_low_high_freq(x):
639
+ fft = torch.fft.fft2(x)
640
+ fft_shifted = torch.fft.fftshift(fft)
641
+ height, width = x.shape[-2:]
642
+ radius = min(height, width) // 5
643
+
644
+ y_grid, x_grid = torch.meshgrid(torch.arange(height), torch.arange(width))
645
+ center_x, center_y = width // 2, height // 2
646
+ mask = (x_grid - center_x) ** 2 + (y_grid - center_y) ** 2 <= radius**2
647
+
648
+ low_freq_mask = mask.unsqueeze(0).unsqueeze(0).to(x.device)
649
+ high_freq_mask = ~low_freq_mask
650
+
651
+ low_freq_fft = fft_shifted * low_freq_mask
652
+ high_freq_fft = fft_shifted * high_freq_mask
653
+
654
+ return low_freq_fft, high_freq_fft
diffusers/hooks/first_block_cache.py ADDED
@@ -0,0 +1,258 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2024 The HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ from dataclasses import dataclass
16
+
17
+ import torch
18
+
19
+ from ..utils import get_logger
20
+ from ..utils.torch_utils import unwrap_module
21
+ from ._common import _ALL_TRANSFORMER_BLOCK_IDENTIFIERS
22
+ from ._helpers import TransformerBlockRegistry
23
+ from .hooks import BaseState, HookRegistry, ModelHook, StateManager
24
+
25
+
26
+ logger = get_logger(__name__) # pylint: disable=invalid-name
27
+
28
+ _FBC_LEADER_BLOCK_HOOK = "fbc_leader_block_hook"
29
+ _FBC_BLOCK_HOOK = "fbc_block_hook"
30
+
31
+
32
+ @dataclass
33
+ class FirstBlockCacheConfig:
34
+ r"""
35
+ Configuration for [First Block
36
+ Cache](https://github.com/chengzeyi/ParaAttention/blob/7a266123671b55e7e5a2fe9af3121f07a36afc78/README.md#first-block-cache-our-dynamic-caching).
37
+
38
+ Args:
39
+ threshold (`float`, defaults to `0.05`):
40
+ The threshold to determine whether or not a forward pass through all layers of the model is required. A
41
+ higher threshold usually results in a forward pass through a lower number of layers and faster inference,
42
+ but might lead to poorer generation quality. A lower threshold may not result in significant generation
43
+ speedup. The threshold is compared against the absmean difference of the residuals between the current and
44
+ cached outputs from the first transformer block. If the difference is below the threshold, the forward pass
45
+ is skipped.
46
+ """
47
+
48
+ threshold: float = 0.05
49
+
50
+
51
+ class FBCSharedBlockState(BaseState):
52
+ def __init__(self) -> None:
53
+ super().__init__()
54
+
55
+ self.head_block_output: torch.Tensor | tuple[torch.Tensor, ...] = None
56
+ self.head_block_residual: torch.Tensor = None
57
+ self.tail_block_residuals: torch.Tensor | tuple[torch.Tensor, ...] = None
58
+ self.should_compute: bool = True
59
+
60
+ def reset(self):
61
+ self.tail_block_residuals = None
62
+ self.should_compute = True
63
+
64
+
65
+ class FBCHeadBlockHook(ModelHook):
66
+ _is_stateful = True
67
+
68
+ def __init__(self, state_manager: StateManager, threshold: float):
69
+ self.state_manager = state_manager
70
+ self.threshold = threshold
71
+ self._metadata = None
72
+
73
+ def initialize_hook(self, module):
74
+ unwrapped_module = unwrap_module(module)
75
+ self._metadata = TransformerBlockRegistry.get(unwrapped_module.__class__)
76
+ return module
77
+
78
+ def new_forward(self, module: torch.nn.Module, *args, **kwargs):
79
+ original_hidden_states = self._metadata._get_parameter_from_args_kwargs("hidden_states", args, kwargs)
80
+
81
+ output = self.fn_ref.original_forward(*args, **kwargs)
82
+ is_output_tuple = isinstance(output, tuple)
83
+
84
+ if is_output_tuple:
85
+ hidden_states_residual = output[self._metadata.return_hidden_states_index] - original_hidden_states
86
+ else:
87
+ hidden_states_residual = output - original_hidden_states
88
+
89
+ shared_state: FBCSharedBlockState = self.state_manager.get_state()
90
+ hidden_states = encoder_hidden_states = None
91
+ should_compute = self._should_compute_remaining_blocks(hidden_states_residual)
92
+ shared_state.should_compute = should_compute
93
+
94
+ if not should_compute:
95
+ # Apply caching
96
+ if is_output_tuple:
97
+ hidden_states = (
98
+ shared_state.tail_block_residuals[0] + output[self._metadata.return_hidden_states_index]
99
+ )
100
+ else:
101
+ hidden_states = shared_state.tail_block_residuals[0] + output
102
+
103
+ if self._metadata.return_encoder_hidden_states_index is not None:
104
+ assert is_output_tuple
105
+ encoder_hidden_states = (
106
+ shared_state.tail_block_residuals[1] + output[self._metadata.return_encoder_hidden_states_index]
107
+ )
108
+
109
+ if is_output_tuple:
110
+ return_output = [None] * len(output)
111
+ return_output[self._metadata.return_hidden_states_index] = hidden_states
112
+ return_output[self._metadata.return_encoder_hidden_states_index] = encoder_hidden_states
113
+ return_output = tuple(return_output)
114
+ else:
115
+ return_output = hidden_states
116
+ output = return_output
117
+ else:
118
+ if is_output_tuple:
119
+ head_block_output = [None] * len(output)
120
+ head_block_output[0] = output[self._metadata.return_hidden_states_index]
121
+ head_block_output[1] = output[self._metadata.return_encoder_hidden_states_index]
122
+ else:
123
+ head_block_output = output
124
+ shared_state.head_block_output = head_block_output
125
+ shared_state.head_block_residual = hidden_states_residual
126
+
127
+ return output
128
+
129
+ def reset_state(self, module):
130
+ self.state_manager.reset()
131
+ return module
132
+
133
+ @torch.compiler.disable
134
+ def _should_compute_remaining_blocks(self, hidden_states_residual: torch.Tensor) -> bool:
135
+ shared_state = self.state_manager.get_state()
136
+ if shared_state.head_block_residual is None:
137
+ return True
138
+ prev_hidden_states_residual = shared_state.head_block_residual
139
+ absmean = (hidden_states_residual - prev_hidden_states_residual).abs().mean()
140
+ prev_hidden_states_absmean = prev_hidden_states_residual.abs().mean()
141
+ diff = (absmean / prev_hidden_states_absmean).item()
142
+ return diff > self.threshold
143
+
144
+
145
+ class FBCBlockHook(ModelHook):
146
+ def __init__(self, state_manager: StateManager, is_tail: bool = False):
147
+ super().__init__()
148
+ self.state_manager = state_manager
149
+ self.is_tail = is_tail
150
+ self._metadata = None
151
+
152
+ def initialize_hook(self, module):
153
+ unwrapped_module = unwrap_module(module)
154
+ self._metadata = TransformerBlockRegistry.get(unwrapped_module.__class__)
155
+ return module
156
+
157
+ def new_forward(self, module: torch.nn.Module, *args, **kwargs):
158
+ original_hidden_states = self._metadata._get_parameter_from_args_kwargs("hidden_states", args, kwargs)
159
+ original_encoder_hidden_states = None
160
+ if self._metadata.return_encoder_hidden_states_index is not None:
161
+ original_encoder_hidden_states = self._metadata._get_parameter_from_args_kwargs(
162
+ "encoder_hidden_states", args, kwargs
163
+ )
164
+
165
+ shared_state = self.state_manager.get_state()
166
+
167
+ if shared_state.should_compute:
168
+ output = self.fn_ref.original_forward(*args, **kwargs)
169
+ if self.is_tail:
170
+ hidden_states_residual = encoder_hidden_states_residual = None
171
+ if isinstance(output, tuple):
172
+ hidden_states_residual = (
173
+ output[self._metadata.return_hidden_states_index] - shared_state.head_block_output[0]
174
+ )
175
+ encoder_hidden_states_residual = (
176
+ output[self._metadata.return_encoder_hidden_states_index] - shared_state.head_block_output[1]
177
+ )
178
+ else:
179
+ hidden_states_residual = output - shared_state.head_block_output
180
+ shared_state.tail_block_residuals = (hidden_states_residual, encoder_hidden_states_residual)
181
+ return output
182
+
183
+ if original_encoder_hidden_states is None:
184
+ return_output = original_hidden_states
185
+ else:
186
+ return_output = [None, None]
187
+ return_output[self._metadata.return_hidden_states_index] = original_hidden_states
188
+ return_output[self._metadata.return_encoder_hidden_states_index] = original_encoder_hidden_states
189
+ return_output = tuple(return_output)
190
+ return return_output
191
+
192
+
193
+ def apply_first_block_cache(module: torch.nn.Module, config: FirstBlockCacheConfig) -> None:
194
+ """
195
+ Applies [First Block
196
+ Cache](https://github.com/chengzeyi/ParaAttention/blob/4de137c5b96416489f06e43e19f2c14a772e28fd/README.md#first-block-cache-our-dynamic-caching)
197
+ to a given module.
198
+
199
+ First Block Cache builds on the ideas of [TeaCache](https://huggingface.co/papers/2411.19108). It is much simpler
200
+ to implement generically for a wide range of models and has been integrated first for experimental purposes.
201
+
202
+ Args:
203
+ module (`torch.nn.Module`):
204
+ The pytorch module to apply FBCache to. Typically, this should be a transformer architecture supported in
205
+ Diffusers, such as `CogVideoXTransformer3DModel`, but external implementations may also work.
206
+ config (`FirstBlockCacheConfig`):
207
+ The configuration to use for applying the FBCache method.
208
+
209
+ Example:
210
+ ```python
211
+ >>> import torch
212
+ >>> from diffusers import CogView4Pipeline
213
+ >>> from diffusers.hooks import apply_first_block_cache, FirstBlockCacheConfig
214
+
215
+ >>> pipe = CogView4Pipeline.from_pretrained("THUDM/CogView4-6B", torch_dtype=torch.bfloat16)
216
+ >>> pipe.to("cuda")
217
+
218
+ >>> apply_first_block_cache(pipe.transformer, FirstBlockCacheConfig(threshold=0.2))
219
+
220
+ >>> prompt = "A photo of an astronaut riding a horse on mars"
221
+ >>> image = pipe(prompt, generator=torch.Generator().manual_seed(42)).images[0]
222
+ >>> image.save("output.png")
223
+ ```
224
+ """
225
+
226
+ state_manager = StateManager(FBCSharedBlockState, (), {})
227
+ remaining_blocks = []
228
+
229
+ for name, submodule in module.named_children():
230
+ if name not in _ALL_TRANSFORMER_BLOCK_IDENTIFIERS or not isinstance(submodule, torch.nn.ModuleList):
231
+ continue
232
+ for index, block in enumerate(submodule):
233
+ remaining_blocks.append((f"{name}.{index}", block))
234
+
235
+ head_block_name, head_block = remaining_blocks.pop(0)
236
+ tail_block_name, tail_block = remaining_blocks.pop(-1)
237
+
238
+ logger.debug(f"Applying FBCHeadBlockHook to '{head_block_name}'")
239
+ _apply_fbc_head_block_hook(head_block, state_manager, config.threshold)
240
+
241
+ for name, block in remaining_blocks:
242
+ logger.debug(f"Applying FBCBlockHook to '{name}'")
243
+ _apply_fbc_block_hook(block, state_manager)
244
+
245
+ logger.debug(f"Applying FBCBlockHook to tail block '{tail_block_name}'")
246
+ _apply_fbc_block_hook(tail_block, state_manager, is_tail=True)
247
+
248
+
249
+ def _apply_fbc_head_block_hook(block: torch.nn.Module, state_manager: StateManager, threshold: float) -> None:
250
+ registry = HookRegistry.check_if_exists_or_initialize(block)
251
+ hook = FBCHeadBlockHook(state_manager, threshold)
252
+ registry.register_hook(hook, _FBC_LEADER_BLOCK_HOOK)
253
+
254
+
255
+ def _apply_fbc_block_hook(block: torch.nn.Module, state_manager: StateManager, is_tail: bool = False) -> None:
256
+ registry = HookRegistry.check_if_exists_or_initialize(block)
257
+ hook = FBCBlockHook(state_manager, is_tail)
258
+ registry.register_hook(hook, _FBC_BLOCK_HOOK)
diffusers/hooks/group_offloading.py ADDED
@@ -0,0 +1,1056 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 The HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ import hashlib
16
+ import os
17
+ from contextlib import contextmanager, nullcontext
18
+ from dataclasses import dataclass, replace
19
+ from enum import Enum
20
+ from typing import Set
21
+
22
+ import safetensors.torch
23
+ import torch
24
+
25
+ from ..utils import get_logger, is_accelerate_available, is_torchao_available
26
+ from ._common import _GO_LC_SUPPORTED_PYTORCH_LAYERS
27
+ from .hooks import HookRegistry, ModelHook
28
+
29
+
30
+ if is_accelerate_available():
31
+ from accelerate.hooks import AlignDevicesHook, CpuOffload
32
+ from accelerate.utils import send_to_device
33
+
34
+
35
+ logger = get_logger(__name__) # pylint: disable=invalid-name
36
+
37
+
38
+ def _is_torchao_tensor(tensor: torch.Tensor) -> bool:
39
+ if not is_torchao_available():
40
+ return False
41
+ from torchao.utils import TorchAOBaseTensor
42
+
43
+ return isinstance(tensor, TorchAOBaseTensor)
44
+
45
+
46
+ def _get_torchao_inner_tensor_names(tensor: torch.Tensor) -> list[str]:
47
+ """Get names of all internal tensor data attributes from a TorchAO tensor."""
48
+ cls = type(tensor)
49
+ names = list(getattr(cls, "tensor_data_names", []))
50
+ for attr_name in getattr(cls, "optional_tensor_data_names", []):
51
+ if getattr(tensor, attr_name, None) is not None:
52
+ names.append(attr_name)
53
+ return names
54
+
55
+
56
+ def _swap_torchao_tensor(param: torch.Tensor, source: torch.Tensor) -> None:
57
+ """Move a TorchAO parameter to the device of `source` via `swap_tensors`.
58
+
59
+ `param.data = source` does not work for `_make_wrapper_subclass` tensors because the `.data` setter only replaces
60
+ the outer wrapper storage while leaving the subclass's internal attributes (e.g. `.qdata`, `.scale`) on the
61
+ original device. `swap_tensors` swaps the full tensor contents in-place, preserving the parameter's identity so
62
+ that any dict keyed by `id(param)` remains valid.
63
+
64
+ Refer to https://github.com/huggingface/diffusers/pull/13276#discussion_r2944471548 for the full discussion.
65
+ """
66
+ torch.utils.swap_tensors(param, source)
67
+
68
+
69
+ def _restore_torchao_tensor(param: torch.Tensor, source: torch.Tensor) -> None:
70
+ """Restore internal tensor data of a TorchAO parameter from `source` without mutating `source`.
71
+
72
+ Unlike `_swap_torchao_tensor` this copies attribute references one-by-one via `setattr` so that `source` is **not**
73
+ modified. Use this when `source` is a cached tensor that must remain unchanged (e.g. a pinned CPU copy in
74
+ `cpu_param_dict`).
75
+ """
76
+ for attr_name in _get_torchao_inner_tensor_names(source):
77
+ setattr(param, attr_name, getattr(source, attr_name))
78
+
79
+
80
+ def _record_stream_torchao_tensor(param: torch.Tensor, stream) -> None:
81
+ """Record stream for all internal tensors of a TorchAO parameter."""
82
+ for attr_name in _get_torchao_inner_tensor_names(param):
83
+ getattr(param, attr_name).record_stream(stream)
84
+
85
+
86
+ # fmt: off
87
+ _GROUP_OFFLOADING = "group_offloading"
88
+ _LAYER_EXECUTION_TRACKER = "layer_execution_tracker"
89
+ _LAZY_PREFETCH_GROUP_OFFLOADING = "lazy_prefetch_group_offloading"
90
+ _GROUP_ID_LAZY_LEAF = "lazy_leafs"
91
+ # fmt: on
92
+
93
+
94
+ class GroupOffloadingType(str, Enum):
95
+ BLOCK_LEVEL = "block_level"
96
+ LEAF_LEVEL = "leaf_level"
97
+
98
+
99
+ @dataclass
100
+ class GroupOffloadingConfig:
101
+ onload_device: torch.device
102
+ offload_device: torch.device
103
+ offload_type: GroupOffloadingType
104
+ non_blocking: bool
105
+ record_stream: bool
106
+ low_cpu_mem_usage: bool
107
+ num_blocks_per_group: int | None = None
108
+ offload_to_disk_path: str | None = None
109
+ stream: torch.cuda.Stream | torch.Stream | None = None
110
+ block_modules: list[str] | None = None
111
+ exclude_kwargs: list[str] | None = None
112
+ module_prefix: str = ""
113
+
114
+
115
+ class ModuleGroup:
116
+ def __init__(
117
+ self,
118
+ modules: list[torch.nn.Module],
119
+ offload_device: torch.device,
120
+ onload_device: torch.device,
121
+ offload_leader: torch.nn.Module,
122
+ onload_leader: torch.nn.Module | None = None,
123
+ parameters: list[torch.nn.Parameter] | None = None,
124
+ buffers: list[torch.Tensor] | None = None,
125
+ non_blocking: bool = False,
126
+ stream: torch.cuda.Stream | torch.Stream | None = None,
127
+ record_stream: bool | None = False,
128
+ low_cpu_mem_usage: bool = False,
129
+ onload_self: bool = True,
130
+ offload_to_disk_path: str | None = None,
131
+ group_id: int | str | None = None,
132
+ ) -> None:
133
+ self.modules = modules
134
+ self.offload_device = offload_device
135
+ self.onload_device = onload_device
136
+ self.offload_leader = offload_leader
137
+ self.onload_leader = onload_leader
138
+ self.parameters = parameters or []
139
+ self.buffers = buffers or []
140
+ self.non_blocking = non_blocking or stream is not None
141
+ self.stream = stream
142
+ self.record_stream = record_stream
143
+ self.onload_self = onload_self
144
+ self.low_cpu_mem_usage = low_cpu_mem_usage
145
+
146
+ self.offload_to_disk_path = offload_to_disk_path
147
+ self._is_offloaded_to_disk = False
148
+
149
+ if self.offload_to_disk_path is not None:
150
+ # Instead of `group_id or str(id(self))` we do this because `group_id` can be "" as well.
151
+ self.group_id = group_id if group_id is not None else str(id(self))
152
+ short_hash = _compute_group_hash(self.group_id)
153
+ self.safetensors_file_path = os.path.join(self.offload_to_disk_path, f"group_{short_hash}.safetensors")
154
+
155
+ all_tensors = []
156
+ for module in self.modules:
157
+ all_tensors.extend(list(module.parameters()))
158
+ all_tensors.extend(list(module.buffers()))
159
+ all_tensors.extend(self.parameters)
160
+ all_tensors.extend(self.buffers)
161
+ all_tensors = list(dict.fromkeys(all_tensors)) # Remove duplicates
162
+
163
+ self.tensor_to_key = {tensor: f"tensor_{i}" for i, tensor in enumerate(all_tensors)}
164
+ self.key_to_tensor = {v: k for k, v in self.tensor_to_key.items()}
165
+ self.cpu_param_dict = {}
166
+ else:
167
+ self.cpu_param_dict = self._init_cpu_param_dict()
168
+
169
+ self._torch_accelerator_module = (
170
+ getattr(torch, torch.accelerator.current_accelerator().type)
171
+ if hasattr(torch, "accelerator")
172
+ else torch.cuda
173
+ )
174
+
175
+ @staticmethod
176
+ def _to_cpu(tensor, low_cpu_mem_usage):
177
+ # For TorchAO tensors, `.data` returns an incomplete wrapper without internal attributes
178
+ # (e.g. `.qdata`, `.scale`), so we must call `.cpu()` on the tensor directly.
179
+ t = tensor.cpu() if _is_torchao_tensor(tensor) else tensor.data.cpu()
180
+ return t if low_cpu_mem_usage else t.pin_memory()
181
+
182
+ def _init_cpu_param_dict(self):
183
+ cpu_param_dict = {}
184
+ if self.stream is None:
185
+ return cpu_param_dict
186
+
187
+ for module in self.modules:
188
+ for param in module.parameters():
189
+ cpu_param_dict[param] = self._to_cpu(param, self.low_cpu_mem_usage)
190
+ for buffer in module.buffers():
191
+ cpu_param_dict[buffer] = self._to_cpu(buffer, self.low_cpu_mem_usage)
192
+
193
+ for param in self.parameters:
194
+ cpu_param_dict[param] = self._to_cpu(param, self.low_cpu_mem_usage)
195
+
196
+ for buffer in self.buffers:
197
+ cpu_param_dict[buffer] = self._to_cpu(buffer, self.low_cpu_mem_usage)
198
+
199
+ return cpu_param_dict
200
+
201
+ @contextmanager
202
+ def _pinned_memory_tensors(self):
203
+ try:
204
+ pinned_dict = {
205
+ param: tensor.pin_memory() if not tensor.is_pinned() else tensor
206
+ for param, tensor in self.cpu_param_dict.items()
207
+ }
208
+ yield pinned_dict
209
+ finally:
210
+ pinned_dict = None
211
+
212
+ def _transfer_tensor_to_device(self, tensor, source_tensor, default_stream):
213
+ moved = source_tensor.to(self.onload_device, non_blocking=self.non_blocking)
214
+ if _is_torchao_tensor(tensor):
215
+ _swap_torchao_tensor(tensor, moved)
216
+ else:
217
+ tensor.data = moved
218
+ if self.record_stream:
219
+ if _is_torchao_tensor(tensor):
220
+ _record_stream_torchao_tensor(tensor, default_stream)
221
+ else:
222
+ tensor.data.record_stream(default_stream)
223
+
224
+ def _process_tensors_from_modules(self, pinned_memory=None, default_stream=None):
225
+ for group_module in self.modules:
226
+ for param in group_module.parameters():
227
+ source = pinned_memory[param] if pinned_memory else param.data
228
+ self._transfer_tensor_to_device(param, source, default_stream)
229
+ for buffer in group_module.buffers():
230
+ source = pinned_memory[buffer] if pinned_memory else buffer.data
231
+ self._transfer_tensor_to_device(buffer, source, default_stream)
232
+
233
+ for param in self.parameters:
234
+ source = pinned_memory[param] if pinned_memory else param.data
235
+ self._transfer_tensor_to_device(param, source, default_stream)
236
+
237
+ for buffer in self.buffers:
238
+ source = pinned_memory[buffer] if pinned_memory else buffer.data
239
+ self._transfer_tensor_to_device(buffer, source, default_stream)
240
+
241
+ def _check_disk_offload_torchao(self):
242
+ all_tensors = list(self.tensor_to_key.keys())
243
+ has_torchao = any(_is_torchao_tensor(t) for t in all_tensors)
244
+ if has_torchao:
245
+ raise ValueError(
246
+ "Disk offloading is not supported for TorchAO quantized tensors because safetensors "
247
+ "cannot serialize TorchAO subclass tensors. Use memory offloading instead by not "
248
+ "setting `offload_to_disk_path`."
249
+ )
250
+
251
+ def _onload_from_disk(self):
252
+ self._check_disk_offload_torchao()
253
+
254
+ if self.stream is not None:
255
+ # Wait for previous Host->Device transfer to complete
256
+ self.stream.synchronize()
257
+
258
+ context = nullcontext() if self.stream is None else self._torch_accelerator_module.stream(self.stream)
259
+ current_stream = self._torch_accelerator_module.current_stream() if self.record_stream else None
260
+
261
+ with context:
262
+ if self.stream is not None:
263
+ # Load to CPU first, pin memory, then async copy to the target device
264
+ loaded_tensors = safetensors.torch.load_file(self.safetensors_file_path, device="cpu")
265
+ for key, tensor_obj in self.key_to_tensor.items():
266
+ pinned_tensor = loaded_tensors[key].pin_memory()
267
+ tensor_obj.data = pinned_tensor.to(self.onload_device, non_blocking=self.non_blocking)
268
+ if self.record_stream:
269
+ tensor_obj.data.record_stream(current_stream)
270
+ else:
271
+ # Load directly to the target device
272
+ onload_device = (
273
+ self.onload_device.type if isinstance(self.onload_device, torch.device) else self.onload_device
274
+ )
275
+ loaded_tensors = safetensors.torch.load_file(self.safetensors_file_path, device=onload_device)
276
+ for key, tensor_obj in self.key_to_tensor.items():
277
+ tensor_obj.data = loaded_tensors[key]
278
+
279
+ def _onload_from_memory(self):
280
+ if self.stream is not None:
281
+ # Wait for previous Host->Device transfer to complete
282
+ self.stream.synchronize()
283
+
284
+ context = nullcontext() if self.stream is None else self._torch_accelerator_module.stream(self.stream)
285
+ default_stream = self._torch_accelerator_module.current_stream() if self.stream is not None else None
286
+
287
+ with context:
288
+ if self.stream is not None:
289
+ with self._pinned_memory_tensors() as pinned_memory:
290
+ self._process_tensors_from_modules(pinned_memory, default_stream=default_stream)
291
+ else:
292
+ self._process_tensors_from_modules(None)
293
+
294
+ def _offload_to_disk(self):
295
+ self._check_disk_offload_torchao()
296
+
297
+ # TODO: we can potentially optimize this code path by checking if the _all_ the desired
298
+ # safetensor files exist on the disk and if so, skip this step entirely, reducing IO
299
+ # overhead. Currently, we just check if the given `safetensors_file_path` exists and if not
300
+ # we perform a write.
301
+ # Check if the file has been saved in this session or if it already exists on disk.
302
+ if not self._is_offloaded_to_disk and not os.path.exists(self.safetensors_file_path):
303
+ os.makedirs(os.path.dirname(self.safetensors_file_path), exist_ok=True)
304
+ tensors_to_save = {key: tensor.data.to(self.offload_device) for tensor, key in self.tensor_to_key.items()}
305
+ safetensors.torch.save_file(tensors_to_save, self.safetensors_file_path)
306
+
307
+ # The group is now considered offloaded to disk for the rest of the session.
308
+ self._is_offloaded_to_disk = True
309
+
310
+ # We do this to free up the RAM which is still holding the up tensor data.
311
+ for tensor_obj in self.tensor_to_key.keys():
312
+ tensor_obj.data = torch.empty_like(tensor_obj.data, device=self.offload_device)
313
+
314
+ def _offload_to_memory(self):
315
+ if self.stream is not None:
316
+ if not self.record_stream:
317
+ self._torch_accelerator_module.current_stream().synchronize()
318
+
319
+ for group_module in self.modules:
320
+ for param in group_module.parameters():
321
+ if _is_torchao_tensor(param):
322
+ _restore_torchao_tensor(param, self.cpu_param_dict[param])
323
+ else:
324
+ param.data = self.cpu_param_dict[param]
325
+ for param in self.parameters:
326
+ if _is_torchao_tensor(param):
327
+ _restore_torchao_tensor(param, self.cpu_param_dict[param])
328
+ else:
329
+ param.data = self.cpu_param_dict[param]
330
+ for buffer in self.buffers:
331
+ if _is_torchao_tensor(buffer):
332
+ _restore_torchao_tensor(buffer, self.cpu_param_dict[buffer])
333
+ else:
334
+ buffer.data = self.cpu_param_dict[buffer]
335
+ else:
336
+ for group_module in self.modules:
337
+ group_module.to(self.offload_device, non_blocking=False)
338
+ for param in self.parameters:
339
+ if _is_torchao_tensor(param):
340
+ moved = param.to(self.offload_device, non_blocking=False)
341
+ _swap_torchao_tensor(param, moved)
342
+ else:
343
+ param.data = param.data.to(self.offload_device, non_blocking=False)
344
+ for buffer in self.buffers:
345
+ if _is_torchao_tensor(buffer):
346
+ moved = buffer.to(self.offload_device, non_blocking=False)
347
+ _swap_torchao_tensor(buffer, moved)
348
+ else:
349
+ buffer.data = buffer.data.to(self.offload_device, non_blocking=False)
350
+
351
+ @torch.compiler.disable()
352
+ def onload_(self):
353
+ r"""Onloads the group of parameters to the onload_device."""
354
+ if self.offload_to_disk_path is not None:
355
+ self._onload_from_disk()
356
+ else:
357
+ self._onload_from_memory()
358
+
359
+ @torch.compiler.disable()
360
+ def offload_(self):
361
+ r"""Offloads the group of parameters to the offload_device."""
362
+ if self.offload_to_disk_path:
363
+ self._offload_to_disk()
364
+ else:
365
+ self._offload_to_memory()
366
+
367
+
368
+ class GroupOffloadingHook(ModelHook):
369
+ r"""
370
+ A hook that offloads groups of torch.nn.Module to the CPU for storage and onloads to accelerator device for
371
+ computation. Each group has one "onload leader" module that is responsible for onloading, and an "offload leader"
372
+ module that is responsible for offloading. If prefetching is enabled, the onload leader of the previous module
373
+ group is responsible for onloading the current module group.
374
+ """
375
+
376
+ _is_stateful = False
377
+
378
+ def __init__(self, group: ModuleGroup, *, config: GroupOffloadingConfig) -> None:
379
+ self.group = group
380
+ self.next_group: ModuleGroup | None = None
381
+ self.config = config
382
+
383
+ def initialize_hook(self, module: torch.nn.Module) -> torch.nn.Module:
384
+ if self.group.offload_leader == module:
385
+ self.group.offload_()
386
+ return module
387
+
388
+ def pre_forward(self, module: torch.nn.Module, *args, **kwargs):
389
+ # If there wasn't an onload_leader assigned, we assume that the submodule that first called its forward
390
+ # method is the onload_leader of the group.
391
+ if self.group.onload_leader is None:
392
+ self.group.onload_leader = module
393
+
394
+ # If the current module is the onload_leader of the group, we onload the group if it is supposed
395
+ # to onload itself. In the case of using prefetching with streams, we onload the next group if
396
+ # it is not supposed to onload itself.
397
+ if self.group.onload_leader == module:
398
+ if self.group.onload_self:
399
+ self.group.onload_()
400
+ else:
401
+ # onload_self=False means this group relies on prefetching from a previous group.
402
+ # However, for conditionally-executed modules (e.g. patch_short/patch_mid/patch_long in Helios),
403
+ # the prefetch chain may not cover them if they were absent during the first forward pass
404
+ # when the execution order was traced. In that case, their weights remain on offload_device,
405
+ # so we fall back to a synchronous onload here.
406
+ params = [p for m in self.group.modules for p in m.parameters()] + list(self.group.parameters)
407
+ if params and params[0].device == self.group.offload_device:
408
+ self.group.onload_()
409
+ if self.group.stream is not None:
410
+ self.group.stream.synchronize()
411
+
412
+ should_onload_next_group = self.next_group is not None and not self.next_group.onload_self
413
+ if should_onload_next_group:
414
+ self.next_group.onload_()
415
+
416
+ should_synchronize = (
417
+ not self.group.onload_self and self.group.stream is not None and not should_onload_next_group
418
+ )
419
+ if should_synchronize:
420
+ # If this group didn't onload itself, it means it was asynchronously onloaded by the
421
+ # previous group. We need to synchronize the side stream to ensure parameters
422
+ # are completely loaded to proceed with forward pass. Without this, uninitialized
423
+ # weights will be used in the computation, leading to incorrect results
424
+ # Also, we should only do this synchronization if we don't already do it from the sync call in
425
+ # self.next_group.onload_, hence the `not should_onload_next_group` check.
426
+ self.group.stream.synchronize()
427
+
428
+ args = send_to_device(args, self.group.onload_device, non_blocking=self.group.non_blocking)
429
+
430
+ # Some Autoencoder models use a feature cache that is passed through submodules
431
+ # and modified in place. The `send_to_device` call returns a copy of this feature cache object
432
+ # which breaks the inplace updates. Use `exclude_kwargs` to mark these cache features
433
+ exclude_kwargs = self.config.exclude_kwargs or []
434
+ if exclude_kwargs:
435
+ moved_kwargs = send_to_device(
436
+ {k: v for k, v in kwargs.items() if k not in exclude_kwargs},
437
+ self.group.onload_device,
438
+ non_blocking=self.group.non_blocking,
439
+ )
440
+ kwargs.update(moved_kwargs)
441
+ else:
442
+ kwargs = send_to_device(kwargs, self.group.onload_device, non_blocking=self.group.non_blocking)
443
+
444
+ return args, kwargs
445
+
446
+ def post_forward(self, module: torch.nn.Module, output):
447
+ if self.group.offload_leader == module:
448
+ self.group.offload_()
449
+ return output
450
+
451
+
452
+ class LazyPrefetchGroupOffloadingHook(ModelHook):
453
+ r"""
454
+ A hook, used in conjunction with GroupOffloadingHook, that applies lazy prefetching to groups of torch.nn.Module.
455
+ This hook is used to determine the order in which the layers are executed during the forward pass. Once the layer
456
+ invocation order is known, assignments of the next_group attribute for prefetching can be made, which allows
457
+ prefetching groups in the correct order.
458
+ """
459
+
460
+ _is_stateful = False
461
+
462
+ def __init__(self):
463
+ self.execution_order: list[tuple[str, torch.nn.Module]] = []
464
+ self._layer_execution_tracker_module_names = set()
465
+
466
+ def initialize_hook(self, module):
467
+ def make_execution_order_update_callback(current_name, current_submodule):
468
+ def callback():
469
+ if not torch.compiler.is_compiling():
470
+ logger.debug(f"Adding {current_name} to the execution order")
471
+ self.execution_order.append((current_name, current_submodule))
472
+
473
+ return callback
474
+
475
+ # To every submodule that contains a group offloading hook (at this point, no prefetching is enabled for any
476
+ # of the groups), we add a layer execution tracker hook that will be used to determine the order in which the
477
+ # layers are executed during the forward pass.
478
+ for name, submodule in module.named_modules():
479
+ if name == "" or not hasattr(submodule, "_diffusers_hook"):
480
+ continue
481
+
482
+ registry = HookRegistry.check_if_exists_or_initialize(submodule)
483
+ group_offloading_hook = registry.get_hook(_GROUP_OFFLOADING)
484
+
485
+ if group_offloading_hook is not None:
486
+ # For the first forward pass, we have to load in a blocking manner
487
+ group_offloading_hook.group.non_blocking = False
488
+ layer_tracker_hook = LayerExecutionTrackerHook(make_execution_order_update_callback(name, submodule))
489
+ registry.register_hook(layer_tracker_hook, _LAYER_EXECUTION_TRACKER)
490
+ self._layer_execution_tracker_module_names.add(name)
491
+
492
+ return module
493
+
494
+ def post_forward(self, module, output):
495
+ # At this point, for the current modules' submodules, we know the execution order of the layers. We can now
496
+ # remove the layer execution tracker hooks and apply prefetching by setting the next_group attribute for each
497
+ # group offloading hook.
498
+ num_executed = len(self.execution_order)
499
+ execution_order_module_names = {name for name, _ in self.execution_order}
500
+
501
+ # It may be possible that some layers were not executed during the forward pass. This can happen if the layer
502
+ # is not used in the forward pass, or if the layer is not executed due to some other reason. In such cases, we
503
+ # may not be able to apply prefetching in the correct order, which can lead to device-mismatch related errors
504
+ # if the missing layers end up being executed in the future.
505
+ if execution_order_module_names != self._layer_execution_tracker_module_names:
506
+ unexecuted_layers = list(self._layer_execution_tracker_module_names - execution_order_module_names)
507
+ if not torch.compiler.is_compiling():
508
+ logger.warning(
509
+ "It seems like some layers were not executed during the forward pass. This may lead to problems when "
510
+ "applying lazy prefetching with automatic tracing and lead to device-mismatch related errors. Please "
511
+ "make sure that all layers are executed during the forward pass. The following layers were not executed:\n"
512
+ f"{unexecuted_layers=}"
513
+ )
514
+
515
+ # Remove the layer execution tracker hooks from the submodules
516
+ base_module_registry = module._diffusers_hook
517
+ registries = [submodule._diffusers_hook for _, submodule in self.execution_order]
518
+ group_offloading_hooks = [registry.get_hook(_GROUP_OFFLOADING) for registry in registries]
519
+
520
+ for i in range(num_executed):
521
+ registries[i].remove_hook(_LAYER_EXECUTION_TRACKER, recurse=False)
522
+
523
+ # Remove the current lazy prefetch group offloading hook so that it doesn't interfere with the next forward pass
524
+ base_module_registry.remove_hook(_LAZY_PREFETCH_GROUP_OFFLOADING, recurse=False)
525
+
526
+ # LazyPrefetchGroupOffloadingHook is only used with streams, so we know that non_blocking should be True.
527
+ # We disable non_blocking for the first forward pass, but need to enable it for the subsequent passes to
528
+ # see the benefits of prefetching.
529
+ for hook in group_offloading_hooks:
530
+ hook.group.non_blocking = True
531
+
532
+ # Set required attributes for prefetching
533
+ if num_executed > 0:
534
+ base_module_group_offloading_hook = base_module_registry.get_hook(_GROUP_OFFLOADING)
535
+ base_module_group_offloading_hook.next_group = group_offloading_hooks[0].group
536
+ base_module_group_offloading_hook.next_group.onload_self = False
537
+
538
+ for i in range(num_executed - 1):
539
+ name1, _ = self.execution_order[i]
540
+ name2, _ = self.execution_order[i + 1]
541
+ if not torch.compiler.is_compiling():
542
+ logger.debug(f"Applying lazy prefetch group offloading from {name1} to {name2}")
543
+ group_offloading_hooks[i].next_group = group_offloading_hooks[i + 1].group
544
+ group_offloading_hooks[i].next_group.onload_self = False
545
+
546
+ return output
547
+
548
+
549
+ class LayerExecutionTrackerHook(ModelHook):
550
+ r"""
551
+ A hook that tracks the order in which the layers are executed during the forward pass by calling back to the
552
+ LazyPrefetchGroupOffloadingHook to update the execution order.
553
+ """
554
+
555
+ _is_stateful = False
556
+
557
+ def __init__(self, execution_order_update_callback):
558
+ self.execution_order_update_callback = execution_order_update_callback
559
+
560
+ def pre_forward(self, module, *args, **kwargs):
561
+ self.execution_order_update_callback()
562
+ return args, kwargs
563
+
564
+
565
+ def apply_group_offloading(
566
+ module: torch.nn.Module,
567
+ onload_device: str | torch.device,
568
+ offload_device: str | torch.device = torch.device("cpu"),
569
+ offload_type: str | GroupOffloadingType = "block_level",
570
+ num_blocks_per_group: int | None = None,
571
+ non_blocking: bool = False,
572
+ use_stream: bool = False,
573
+ record_stream: bool = False,
574
+ low_cpu_mem_usage: bool = False,
575
+ offload_to_disk_path: str | None = None,
576
+ block_modules: list[str] | None = None,
577
+ exclude_kwargs: list[str] | None = None,
578
+ ) -> None:
579
+ r"""
580
+ Applies group offloading to the internal layers of a torch.nn.Module. To understand what group offloading is, and
581
+ where it is beneficial, we need to first provide some context on how other supported offloading methods work.
582
+
583
+ Typically, offloading is done at two levels:
584
+ - Module-level: In Diffusers, this can be enabled using the `ModelMixin::enable_model_cpu_offload()` method. It
585
+ works by offloading each component of a pipeline to the CPU for storage, and onloading to the accelerator device
586
+ when needed for computation. This method is more memory-efficient than keeping all components on the accelerator,
587
+ but the memory requirements are still quite high. For this method to work, one needs memory equivalent to size of
588
+ the model in runtime dtype + size of largest intermediate activation tensors to be able to complete the forward
589
+ pass.
590
+ - Leaf-level: In Diffusers, this can be enabled using the `ModelMixin::enable_sequential_cpu_offload()` method. It
591
+ works by offloading the lowest leaf-level parameters of the computation graph to the CPU for storage, and
592
+ onloading only the leafs to the accelerator device for computation. This uses the lowest amount of accelerator
593
+ memory, but can be slower due to the excessive number of device synchronizations.
594
+
595
+ Group offloading is a middle ground between the two methods. It works by offloading groups of internal layers,
596
+ (either `torch.nn.ModuleList` or `torch.nn.Sequential`). This method uses lower memory than module-level
597
+ offloading. It is also faster than leaf-level/sequential offloading, as the number of device synchronizations is
598
+ reduced.
599
+
600
+ Another supported feature (for CUDA devices with support for asynchronous data transfer streams) is the ability to
601
+ overlap data transfer and computation to reduce the overall execution time compared to sequential offloading. This
602
+ is enabled using layer prefetching with streams, i.e., the layer that is to be executed next starts onloading to
603
+ the accelerator device while the current layer is being executed - this increases the memory requirements slightly.
604
+ Note that this implementation also supports leaf-level offloading but can be made much faster when using streams.
605
+
606
+ Args:
607
+ module (`torch.nn.Module`):
608
+ The module to which group offloading is applied.
609
+ onload_device (`torch.device`):
610
+ The device to which the group of modules are onloaded.
611
+ offload_device (`torch.device`, defaults to `torch.device("cpu")`):
612
+ The device to which the group of modules are offloaded. This should typically be the CPU. Default is CPU.
613
+ offload_type (`str` or `GroupOffloadingType`, defaults to "block_level"):
614
+ The type of offloading to be applied. Can be one of "block_level" or "leaf_level". Default is
615
+ "block_level".
616
+ offload_to_disk_path (`str`, *optional*, defaults to `None`):
617
+ The path to the directory where parameters will be offloaded. Setting this option can be useful in limited
618
+ RAM environment settings where a reasonable speed-memory trade-off is desired.
619
+ num_blocks_per_group (`int`, *optional*):
620
+ The number of blocks per group when using offload_type="block_level". This is required when using
621
+ offload_type="block_level".
622
+ non_blocking (`bool`, defaults to `False`):
623
+ If True, offloading and onloading is done with non-blocking data transfer.
624
+ use_stream (`bool`, defaults to `False`):
625
+ If True, offloading and onloading is done asynchronously using a CUDA stream. This can be useful for
626
+ overlapping computation and data transfer.
627
+ record_stream (`bool`, defaults to `False`): When enabled with `use_stream`, it marks the current tensor
628
+ as having been used by this stream. It is faster at the expense of slightly more memory usage. Refer to the
629
+ [PyTorch official docs](https://pytorch.org/docs/stable/generated/torch.Tensor.record_stream.html) more
630
+ details.
631
+ low_cpu_mem_usage (`bool`, defaults to `False`):
632
+ If True, the CPU memory usage is minimized by pinning tensors on-the-fly instead of pre-pinning them. This
633
+ option only matters when using streamed CPU offloading (i.e. `use_stream=True`). This can be useful when
634
+ the CPU memory is a bottleneck but may counteract the benefits of using streams.
635
+ block_modules (`list[str]`, *optional*):
636
+ List of module names that should be treated as blocks for offloading. If provided, only these modules will
637
+ be considered for block-level offloading. If not provided, the default block detection logic will be used.
638
+ exclude_kwargs (`list[str]`, *optional*):
639
+ List of kwarg keys that should not be processed by send_to_device. This is useful for mutable state like
640
+ caching lists that need to maintain their object identity across forward passes. If not provided, will be
641
+ inferred from the module's `_skip_keys` attribute if it exists.
642
+
643
+ Example:
644
+ ```python
645
+ >>> from diffusers import CogVideoXTransformer3DModel
646
+ >>> from diffusers.hooks import apply_group_offloading
647
+
648
+ >>> transformer = CogVideoXTransformer3DModel.from_pretrained(
649
+ ... "THUDM/CogVideoX-5b", subfolder="transformer", torch_dtype=torch.bfloat16
650
+ ... )
651
+
652
+ >>> apply_group_offloading(
653
+ ... transformer,
654
+ ... onload_device=torch.device("cuda"),
655
+ ... offload_device=torch.device("cpu"),
656
+ ... offload_type="block_level",
657
+ ... num_blocks_per_group=2,
658
+ ... use_stream=True,
659
+ ... )
660
+ ```
661
+ """
662
+
663
+ onload_device = torch.device(onload_device) if isinstance(onload_device, str) else onload_device
664
+ offload_device = torch.device(offload_device) if isinstance(offload_device, str) else offload_device
665
+ offload_type = GroupOffloadingType(offload_type)
666
+
667
+ stream = None
668
+ if use_stream:
669
+ if torch.cuda.is_available():
670
+ stream = torch.cuda.Stream()
671
+ elif hasattr(torch, "xpu") and torch.xpu.is_available():
672
+ stream = torch.Stream()
673
+ else:
674
+ raise ValueError("Using streams for data transfer requires a CUDA device, or an Intel XPU device.")
675
+
676
+ if not use_stream and record_stream:
677
+ raise ValueError("`record_stream` cannot be True when `use_stream=False`.")
678
+ if offload_type == GroupOffloadingType.BLOCK_LEVEL and num_blocks_per_group is None:
679
+ raise ValueError("`num_blocks_per_group` must be provided when using `offload_type='block_level'.")
680
+
681
+ _raise_error_if_accelerate_model_or_sequential_hook_present(module)
682
+
683
+ if block_modules is None:
684
+ block_modules = getattr(module, "_group_offload_block_modules", None)
685
+
686
+ if exclude_kwargs is None:
687
+ exclude_kwargs = getattr(module, "_skip_keys", None)
688
+
689
+ config = GroupOffloadingConfig(
690
+ onload_device=onload_device,
691
+ offload_device=offload_device,
692
+ offload_type=offload_type,
693
+ num_blocks_per_group=num_blocks_per_group,
694
+ non_blocking=non_blocking,
695
+ stream=stream,
696
+ record_stream=record_stream,
697
+ low_cpu_mem_usage=low_cpu_mem_usage,
698
+ offload_to_disk_path=offload_to_disk_path,
699
+ block_modules=block_modules,
700
+ exclude_kwargs=exclude_kwargs,
701
+ )
702
+ _apply_group_offloading(module, config)
703
+
704
+
705
+ def _apply_group_offloading(module: torch.nn.Module, config: GroupOffloadingConfig) -> None:
706
+ if config.offload_type == GroupOffloadingType.BLOCK_LEVEL:
707
+ _apply_group_offloading_block_level(module, config)
708
+ elif config.offload_type == GroupOffloadingType.LEAF_LEVEL:
709
+ _apply_group_offloading_leaf_level(module, config)
710
+ else:
711
+ assert False
712
+
713
+
714
+ def _apply_group_offloading_block_level(module: torch.nn.Module, config: GroupOffloadingConfig) -> None:
715
+ r"""
716
+ This function applies offloading to groups of torch.nn.ModuleList or torch.nn.Sequential blocks, and explicitly
717
+ defined block modules. In comparison to the "leaf_level" offloading, which is more fine-grained, this offloading is
718
+ done at the top-level blocks and modules specified in block_modules.
719
+
720
+ When block_modules is provided, only those modules will be treated as blocks for offloading. For each specified
721
+ module, recursively apply block offloading to it.
722
+ """
723
+ if config.stream is not None and config.num_blocks_per_group != 1:
724
+ logger.warning(
725
+ f"Using streams is only supported for num_blocks_per_group=1. Got {config.num_blocks_per_group=}. Setting it to 1."
726
+ )
727
+ config.num_blocks_per_group = 1
728
+
729
+ block_modules = set(config.block_modules) if config.block_modules is not None else set()
730
+
731
+ # Create module groups for ModuleList and Sequential blocks, and explicitly defined block modules
732
+ modules_with_group_offloading = set()
733
+ unmatched_modules = []
734
+ matched_module_groups = []
735
+
736
+ for name, submodule in module.named_children():
737
+ # Check if this is an explicitly defined block module
738
+ if name in block_modules:
739
+ # Track submodule using a prefix to avoid filename collisions during disk offload.
740
+ # Without this, submodules sharing the same model class would be assigned identical
741
+ # filenames (derived from the class name).
742
+ prefix = f"{config.module_prefix}{name}." if config.module_prefix else f"{name}."
743
+ submodule_config = replace(config, module_prefix=prefix)
744
+
745
+ _apply_group_offloading_block_level(submodule, submodule_config)
746
+ modules_with_group_offloading.add(name)
747
+
748
+ elif isinstance(submodule, (torch.nn.ModuleList, torch.nn.Sequential)):
749
+ # Handle ModuleList and Sequential blocks as before
750
+ for i in range(0, len(submodule), config.num_blocks_per_group):
751
+ current_modules = list(submodule[i : i + config.num_blocks_per_group])
752
+ if len(current_modules) == 0:
753
+ continue
754
+
755
+ group_id = f"{config.module_prefix}{name}_{i}_{i + len(current_modules) - 1}"
756
+ group = ModuleGroup(
757
+ modules=current_modules,
758
+ offload_device=config.offload_device,
759
+ onload_device=config.onload_device,
760
+ offload_to_disk_path=config.offload_to_disk_path,
761
+ offload_leader=current_modules[-1],
762
+ onload_leader=current_modules[0],
763
+ non_blocking=config.non_blocking,
764
+ stream=config.stream,
765
+ record_stream=config.record_stream,
766
+ low_cpu_mem_usage=config.low_cpu_mem_usage,
767
+ onload_self=True,
768
+ group_id=group_id,
769
+ )
770
+ matched_module_groups.append(group)
771
+ for j in range(i, i + len(current_modules)):
772
+ modules_with_group_offloading.add(f"{name}.{j}")
773
+ else:
774
+ # This is an unmatched module
775
+ unmatched_modules.append((name, submodule))
776
+
777
+ # Apply group offloading hooks to the module groups
778
+ for i, group in enumerate(matched_module_groups):
779
+ for group_module in group.modules:
780
+ _apply_group_offloading_hook(group_module, group, config=config)
781
+
782
+ # Parameters and Buffers of the top-level module need to be offloaded/onloaded separately
783
+ # when the forward pass of this module is called. This is because the top-level module is not
784
+ # part of any group (as doing so would lead to no VRAM savings).
785
+ parameters = _gather_parameters_with_no_group_offloading_parent(module, modules_with_group_offloading)
786
+ buffers = _gather_buffers_with_no_group_offloading_parent(module, modules_with_group_offloading)
787
+ parameters = [param for _, param in parameters]
788
+ buffers = [buffer for _, buffer in buffers]
789
+
790
+ # Create a group for the remaining unmatched submodules of the top-level
791
+ # module so that they are on the correct device when the forward pass is called.
792
+ unmatched_modules = [unmatched_module for _, unmatched_module in unmatched_modules]
793
+ if len(unmatched_modules) > 0 or len(parameters) > 0 or len(buffers) > 0:
794
+ unmatched_group = ModuleGroup(
795
+ modules=unmatched_modules,
796
+ offload_device=config.offload_device,
797
+ onload_device=config.onload_device,
798
+ offload_to_disk_path=config.offload_to_disk_path,
799
+ offload_leader=module,
800
+ onload_leader=module,
801
+ parameters=parameters,
802
+ buffers=buffers,
803
+ non_blocking=False,
804
+ stream=None,
805
+ record_stream=False,
806
+ onload_self=True,
807
+ group_id=f"{config.module_prefix}{module.__class__.__name__}_unmatched_group",
808
+ )
809
+ if config.stream is None:
810
+ _apply_group_offloading_hook(module, unmatched_group, config=config)
811
+ else:
812
+ _apply_lazy_group_offloading_hook(module, unmatched_group, config=config)
813
+
814
+
815
+ def _apply_group_offloading_leaf_level(module: torch.nn.Module, config: GroupOffloadingConfig) -> None:
816
+ r"""
817
+ This function applies offloading to groups of leaf modules in a torch.nn.Module. This method has minimal memory
818
+ requirements. However, it can be slower compared to other offloading methods due to the excessive number of device
819
+ synchronizations. When using devices that support streams to overlap data transfer and computation, this method can
820
+ reduce memory usage without any performance degradation.
821
+ """
822
+ # Create module groups for leaf modules and apply group offloading hooks
823
+ modules_with_group_offloading = set()
824
+ for name, submodule in module.named_modules():
825
+ if not isinstance(submodule, _GO_LC_SUPPORTED_PYTORCH_LAYERS):
826
+ continue
827
+ group = ModuleGroup(
828
+ modules=[submodule],
829
+ offload_device=config.offload_device,
830
+ onload_device=config.onload_device,
831
+ offload_to_disk_path=config.offload_to_disk_path,
832
+ offload_leader=submodule,
833
+ onload_leader=submodule,
834
+ non_blocking=config.non_blocking,
835
+ stream=config.stream,
836
+ record_stream=config.record_stream,
837
+ low_cpu_mem_usage=config.low_cpu_mem_usage,
838
+ onload_self=True,
839
+ group_id=name,
840
+ )
841
+ _apply_group_offloading_hook(submodule, group, config=config)
842
+ modules_with_group_offloading.add(name)
843
+
844
+ # Parameters and Buffers at all non-leaf levels need to be offloaded/onloaded separately when the forward pass
845
+ # of the module is called
846
+ module_dict = dict(module.named_modules())
847
+ parameters = _gather_parameters_with_no_group_offloading_parent(module, modules_with_group_offloading)
848
+ buffers = _gather_buffers_with_no_group_offloading_parent(module, modules_with_group_offloading)
849
+
850
+ # Find closest module parent for each parameter and buffer, and attach group hooks
851
+ parent_to_parameters = {}
852
+ for name, param in parameters:
853
+ parent_name = _find_parent_module_in_module_dict(name, module_dict)
854
+ if parent_name in parent_to_parameters:
855
+ parent_to_parameters[parent_name].append(param)
856
+ else:
857
+ parent_to_parameters[parent_name] = [param]
858
+
859
+ parent_to_buffers = {}
860
+ for name, buffer in buffers:
861
+ parent_name = _find_parent_module_in_module_dict(name, module_dict)
862
+ if parent_name in parent_to_buffers:
863
+ parent_to_buffers[parent_name].append(buffer)
864
+ else:
865
+ parent_to_buffers[parent_name] = [buffer]
866
+
867
+ parent_names = set(parent_to_parameters.keys()) | set(parent_to_buffers.keys())
868
+ for name in parent_names:
869
+ parameters = parent_to_parameters.get(name, [])
870
+ buffers = parent_to_buffers.get(name, [])
871
+ parent_module = module_dict[name]
872
+ group = ModuleGroup(
873
+ modules=[],
874
+ offload_device=config.offload_device,
875
+ onload_device=config.onload_device,
876
+ offload_leader=parent_module,
877
+ onload_leader=parent_module,
878
+ offload_to_disk_path=config.offload_to_disk_path,
879
+ parameters=parameters,
880
+ buffers=buffers,
881
+ non_blocking=config.non_blocking,
882
+ stream=config.stream,
883
+ record_stream=config.record_stream,
884
+ low_cpu_mem_usage=config.low_cpu_mem_usage,
885
+ onload_self=True,
886
+ group_id=name,
887
+ )
888
+ _apply_group_offloading_hook(parent_module, group, config=config)
889
+
890
+ if config.stream is not None:
891
+ # When using streams, we need to know the layer execution order for applying prefetching (to overlap data transfer
892
+ # and computation). Since we don't know the order beforehand, we apply a lazy prefetching hook that will find the
893
+ # execution order and apply prefetching in the correct order.
894
+ unmatched_group = ModuleGroup(
895
+ modules=[],
896
+ offload_device=config.offload_device,
897
+ onload_device=config.onload_device,
898
+ offload_to_disk_path=config.offload_to_disk_path,
899
+ offload_leader=module,
900
+ onload_leader=module,
901
+ parameters=None,
902
+ buffers=None,
903
+ non_blocking=False,
904
+ stream=None,
905
+ record_stream=False,
906
+ low_cpu_mem_usage=config.low_cpu_mem_usage,
907
+ onload_self=True,
908
+ group_id=_GROUP_ID_LAZY_LEAF,
909
+ )
910
+ _apply_lazy_group_offloading_hook(module, unmatched_group, config=config)
911
+
912
+
913
+ def _apply_group_offloading_hook(
914
+ module: torch.nn.Module,
915
+ group: ModuleGroup,
916
+ *,
917
+ config: GroupOffloadingConfig,
918
+ ) -> None:
919
+ registry = HookRegistry.check_if_exists_or_initialize(module)
920
+
921
+ # We may have already registered a group offloading hook if the module had a torch.nn.Parameter whose parent
922
+ # is the current module. In such cases, we don't want to overwrite the existing group offloading hook.
923
+ if registry.get_hook(_GROUP_OFFLOADING) is None:
924
+ hook = GroupOffloadingHook(group, config=config)
925
+ registry.register_hook(hook, _GROUP_OFFLOADING)
926
+
927
+
928
+ def _apply_lazy_group_offloading_hook(
929
+ module: torch.nn.Module,
930
+ group: ModuleGroup,
931
+ *,
932
+ config: GroupOffloadingConfig,
933
+ ) -> None:
934
+ registry = HookRegistry.check_if_exists_or_initialize(module)
935
+
936
+ # We may have already registered a group offloading hook if the module had a torch.nn.Parameter whose parent
937
+ # is the current module. In such cases, we don't want to overwrite the existing group offloading hook.
938
+ if registry.get_hook(_GROUP_OFFLOADING) is None:
939
+ hook = GroupOffloadingHook(group, config=config)
940
+ registry.register_hook(hook, _GROUP_OFFLOADING)
941
+
942
+ lazy_prefetch_hook = LazyPrefetchGroupOffloadingHook()
943
+ registry.register_hook(lazy_prefetch_hook, _LAZY_PREFETCH_GROUP_OFFLOADING)
944
+
945
+
946
+ def _gather_parameters_with_no_group_offloading_parent(
947
+ module: torch.nn.Module, modules_with_group_offloading: Set[str]
948
+ ) -> list[torch.nn.Parameter]:
949
+ parameters = []
950
+ for name, parameter in module.named_parameters():
951
+ has_parent_with_group_offloading = False
952
+ atoms = name.split(".")
953
+ while len(atoms) > 0:
954
+ parent_name = ".".join(atoms)
955
+ if parent_name in modules_with_group_offloading:
956
+ has_parent_with_group_offloading = True
957
+ break
958
+ atoms.pop()
959
+ if not has_parent_with_group_offloading:
960
+ parameters.append((name, parameter))
961
+ return parameters
962
+
963
+
964
+ def _gather_buffers_with_no_group_offloading_parent(
965
+ module: torch.nn.Module, modules_with_group_offloading: Set[str]
966
+ ) -> list[torch.Tensor]:
967
+ buffers = []
968
+ for name, buffer in module.named_buffers():
969
+ has_parent_with_group_offloading = False
970
+ atoms = name.split(".")
971
+ while len(atoms) > 0:
972
+ parent_name = ".".join(atoms)
973
+ if parent_name in modules_with_group_offloading:
974
+ has_parent_with_group_offloading = True
975
+ break
976
+ atoms.pop()
977
+ if not has_parent_with_group_offloading:
978
+ buffers.append((name, buffer))
979
+ return buffers
980
+
981
+
982
+ def _find_parent_module_in_module_dict(name: str, module_dict: dict[str, torch.nn.Module]) -> str:
983
+ atoms = name.split(".")
984
+ while len(atoms) > 0:
985
+ parent_name = ".".join(atoms)
986
+ if parent_name in module_dict:
987
+ return parent_name
988
+ atoms.pop()
989
+ return ""
990
+
991
+
992
+ def _raise_error_if_accelerate_model_or_sequential_hook_present(module: torch.nn.Module) -> None:
993
+ if not is_accelerate_available():
994
+ return
995
+ for name, submodule in module.named_modules():
996
+ if not hasattr(submodule, "_hf_hook"):
997
+ continue
998
+ if isinstance(submodule._hf_hook, (AlignDevicesHook, CpuOffload)):
999
+ raise ValueError(
1000
+ f"Cannot apply group offloading to a module that is already applying an alternative "
1001
+ f"offloading strategy from Accelerate. If you want to apply group offloading, please "
1002
+ f"disable the existing offloading strategy first. Offending module: {name} ({type(submodule)})"
1003
+ )
1004
+
1005
+
1006
+ def _get_top_level_group_offload_hook(module: torch.nn.Module) -> GroupOffloadingHook | None:
1007
+ for submodule in module.modules():
1008
+ if hasattr(submodule, "_diffusers_hook"):
1009
+ group_offloading_hook = submodule._diffusers_hook.get_hook(_GROUP_OFFLOADING)
1010
+ if group_offloading_hook is not None:
1011
+ return group_offloading_hook
1012
+ return None
1013
+
1014
+
1015
+ def _is_group_offload_enabled(module: torch.nn.Module) -> bool:
1016
+ top_level_group_offload_hook = _get_top_level_group_offload_hook(module)
1017
+ return top_level_group_offload_hook is not None
1018
+
1019
+
1020
+ def _get_group_onload_device(module: torch.nn.Module) -> torch.device:
1021
+ top_level_group_offload_hook = _get_top_level_group_offload_hook(module)
1022
+ if top_level_group_offload_hook is not None:
1023
+ return top_level_group_offload_hook.config.onload_device
1024
+ raise ValueError("Group offloading is not enabled for the provided module.")
1025
+
1026
+
1027
+ def _compute_group_hash(group_id):
1028
+ hashed_id = hashlib.sha256(group_id.encode("utf-8")).hexdigest()
1029
+ # first 16 characters for a reasonably short but unique name
1030
+ return hashed_id[:16]
1031
+
1032
+
1033
+ def _maybe_remove_and_reapply_group_offloading(module: torch.nn.Module) -> None:
1034
+ r"""
1035
+ Removes the group offloading hook from the module and re-applies it. This is useful when the module has been
1036
+ modified in-place and the group offloading hook references-to-tensors needs to be updated. The in-place
1037
+ modification can happen in a number of ways, for example, fusing QKV or unloading/loading LoRAs on-the-fly.
1038
+
1039
+ In this implementation, we make an assumption that group offloading has only been applied at the top-level module,
1040
+ and therefore all submodules have the same onload and offload devices. If this assumption is not true, say in the
1041
+ case where user has applied group offloading at multiple levels, this function will not work as expected.
1042
+
1043
+ There is some performance penalty associated with doing this when non-default streams are used, because we need to
1044
+ retrace the execution order of the layers with `LazyPrefetchGroupOffloadingHook`.
1045
+ """
1046
+ top_level_group_offload_hook = _get_top_level_group_offload_hook(module)
1047
+
1048
+ if top_level_group_offload_hook is None:
1049
+ return
1050
+
1051
+ registry = HookRegistry.check_if_exists_or_initialize(module)
1052
+ registry.remove_hook(_GROUP_OFFLOADING, recurse=True)
1053
+ registry.remove_hook(_LAYER_EXECUTION_TRACKER, recurse=True)
1054
+ registry.remove_hook(_LAZY_PREFETCH_GROUP_OFFLOADING, recurse=True)
1055
+
1056
+ _apply_group_offloading(module, top_level_group_offload_hook.config)
diffusers/hooks/hooks.py ADDED
@@ -0,0 +1,312 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 The HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ import functools
16
+ from typing import Any
17
+
18
+ import torch
19
+
20
+ from ..utils.logging import get_logger
21
+ from ..utils.torch_utils import unwrap_module
22
+
23
+
24
+ logger = get_logger(__name__) # pylint: disable=invalid-name
25
+
26
+
27
+ class BaseState:
28
+ def reset(self, *args, **kwargs) -> None:
29
+ raise NotImplementedError(
30
+ "BaseState::reset is not implemented. Please implement this method in the derived class."
31
+ )
32
+
33
+
34
+ class StateManager:
35
+ def __init__(self, state_cls: BaseState, init_args=None, init_kwargs=None):
36
+ self._state_cls = state_cls
37
+ self._init_args = init_args if init_args is not None else ()
38
+ self._init_kwargs = init_kwargs if init_kwargs is not None else {}
39
+ self._state_cache = {}
40
+ self._current_context = None
41
+
42
+ def get_state(self):
43
+ if self._current_context is None:
44
+ raise ValueError("No context is set. Please set a context before retrieving the state.")
45
+ if self._current_context not in self._state_cache.keys():
46
+ self._state_cache[self._current_context] = self._state_cls(*self._init_args, **self._init_kwargs)
47
+ return self._state_cache[self._current_context]
48
+
49
+ def set_context(self, name: str) -> None:
50
+ self._current_context = name
51
+
52
+ def reset(self, *args, **kwargs) -> None:
53
+ for name, state in list(self._state_cache.items()):
54
+ state.reset(*args, **kwargs)
55
+ self._state_cache.pop(name)
56
+ self._current_context = None
57
+
58
+
59
+ class ModelHook:
60
+ r"""
61
+ A hook that contains callbacks to be executed just before and after the forward method of a model.
62
+ """
63
+
64
+ _is_stateful = False
65
+
66
+ def __init__(self):
67
+ self.fn_ref: "HookFunctionReference" = None
68
+
69
+ def initialize_hook(self, module: torch.nn.Module) -> torch.nn.Module:
70
+ r"""
71
+ Hook that is executed when a model is initialized.
72
+
73
+ Args:
74
+ module (`torch.nn.Module`):
75
+ The module attached to this hook.
76
+ """
77
+ return module
78
+
79
+ def deinitalize_hook(self, module: torch.nn.Module) -> torch.nn.Module:
80
+ r"""
81
+ Hook that is executed when a model is deinitialized.
82
+
83
+ Args:
84
+ module (`torch.nn.Module`):
85
+ The module attached to this hook.
86
+ """
87
+ return module
88
+
89
+ def pre_forward(self, module: torch.nn.Module, *args, **kwargs) -> tuple[tuple[Any], dict[str, Any]]:
90
+ r"""
91
+ Hook that is executed just before the forward method of the model.
92
+
93
+ Args:
94
+ module (`torch.nn.Module`):
95
+ The module whose forward pass will be executed just after this event.
96
+ args (`tuple[Any]`):
97
+ The positional arguments passed to the module.
98
+ kwargs (`dict[Str, Any]`):
99
+ The keyword arguments passed to the module.
100
+ Returns:
101
+ `tuple[tuple[Any], dict[Str, Any]]`:
102
+ A tuple with the treated `args` and `kwargs`.
103
+ """
104
+ return args, kwargs
105
+
106
+ def post_forward(self, module: torch.nn.Module, output: Any) -> Any:
107
+ r"""
108
+ Hook that is executed just after the forward method of the model.
109
+
110
+ Args:
111
+ module (`torch.nn.Module`):
112
+ The module whose forward pass been executed just before this event.
113
+ output (`Any`):
114
+ The output of the module.
115
+ Returns:
116
+ `Any`: The processed `output`.
117
+ """
118
+ return output
119
+
120
+ def detach_hook(self, module: torch.nn.Module) -> torch.nn.Module:
121
+ r"""
122
+ Hook that is executed when the hook is detached from a module.
123
+
124
+ Args:
125
+ module (`torch.nn.Module`):
126
+ The module detached from this hook.
127
+ """
128
+ return module
129
+
130
+ def reset_state(self, module: torch.nn.Module):
131
+ if self._is_stateful:
132
+ raise NotImplementedError("This hook is stateful and needs to implement the `reset_state` method.")
133
+ return module
134
+
135
+ def _set_context(self, module: torch.nn.Module, name: str) -> None:
136
+ # Iterate over all attributes of the hook to see if any of them have the type `StateManager`. If so, call `set_context` on them.
137
+ for attr_name in dir(self):
138
+ attr = getattr(self, attr_name)
139
+ if isinstance(attr, StateManager):
140
+ attr.set_context(name)
141
+ return module
142
+
143
+
144
+ class HookFunctionReference:
145
+ def __init__(self) -> None:
146
+ """A container class that maintains mutable references to forward pass functions in a hook chain.
147
+
148
+ Its mutable nature allows the hook system to modify the execution chain dynamically without rebuilding the
149
+ entire forward pass structure.
150
+
151
+ Attributes:
152
+ pre_forward: A callable that processes inputs before the main forward pass.
153
+ post_forward: A callable that processes outputs after the main forward pass.
154
+ forward: The current forward function in the hook chain.
155
+ original_forward: The original forward function, stored when a hook provides a custom new_forward.
156
+
157
+ The class enables hook removal by allowing updates to the forward chain through reference modification rather
158
+ than requiring reconstruction of the entire chain. When a hook is removed, only the relevant references need to
159
+ be updated, preserving the execution order of the remaining hooks.
160
+ """
161
+ self.pre_forward = None
162
+ self.post_forward = None
163
+ self.forward = None
164
+ self.original_forward = None
165
+
166
+
167
+ class HookRegistry:
168
+ def __init__(self, module_ref: torch.nn.Module) -> None:
169
+ super().__init__()
170
+
171
+ self.hooks: dict[str, ModelHook] = {}
172
+
173
+ self._module_ref = module_ref
174
+ self._hook_order = []
175
+ self._fn_refs = []
176
+
177
+ def register_hook(self, hook: ModelHook, name: str) -> None:
178
+ if name in self.hooks.keys():
179
+ raise ValueError(
180
+ f"Hook with name {name} already exists in the registry. Please use a different name or "
181
+ f"first remove the existing hook and then add a new one."
182
+ )
183
+
184
+ self._module_ref = hook.initialize_hook(self._module_ref)
185
+
186
+ def create_new_forward(function_reference: HookFunctionReference):
187
+ def new_forward(module, *args, **kwargs):
188
+ args, kwargs = function_reference.pre_forward(module, *args, **kwargs)
189
+ output = function_reference.forward(*args, **kwargs)
190
+ return function_reference.post_forward(module, output)
191
+
192
+ return new_forward
193
+
194
+ forward = self._module_ref.forward
195
+
196
+ fn_ref = HookFunctionReference()
197
+ fn_ref.pre_forward = hook.pre_forward
198
+ fn_ref.post_forward = hook.post_forward
199
+ fn_ref.forward = forward
200
+
201
+ if hasattr(hook, "new_forward"):
202
+ fn_ref.original_forward = forward
203
+ fn_ref.forward = functools.update_wrapper(
204
+ functools.partial(hook.new_forward, self._module_ref), hook.new_forward
205
+ )
206
+
207
+ rewritten_forward = create_new_forward(fn_ref)
208
+ # Wrap from the original `forward` so `inspect.signature` follows `__wrapped__` to the real
209
+ # signature instead of the generic `(module, *args, **kwargs)`, which breaks `torch.export`.
210
+ self._module_ref.forward = functools.update_wrapper(
211
+ functools.partial(rewritten_forward, self._module_ref), forward
212
+ )
213
+
214
+ hook.fn_ref = fn_ref
215
+ self.hooks[name] = hook
216
+ self._hook_order.append(name)
217
+ self._fn_refs.append(fn_ref)
218
+
219
+ def get_hook(self, name: str) -> ModelHook | None:
220
+ return self.hooks.get(name, None)
221
+
222
+ def remove_hook(self, name: str, recurse: bool = True) -> None:
223
+ if name in self.hooks.keys():
224
+ num_hooks = len(self._hook_order)
225
+ hook = self.hooks[name]
226
+ index = self._hook_order.index(name)
227
+ fn_ref = self._fn_refs[index]
228
+
229
+ old_forward = fn_ref.forward
230
+ if fn_ref.original_forward is not None:
231
+ old_forward = fn_ref.original_forward
232
+
233
+ if index == num_hooks - 1:
234
+ self._module_ref.forward = old_forward
235
+ else:
236
+ self._fn_refs[index + 1].forward = old_forward
237
+
238
+ self._module_ref = hook.deinitalize_hook(self._module_ref)
239
+ del self.hooks[name]
240
+ self._hook_order.pop(index)
241
+ self._fn_refs.pop(index)
242
+
243
+ if recurse:
244
+ for module_name, module in self._module_ref.named_modules():
245
+ if module_name == "":
246
+ continue
247
+ if hasattr(module, "_diffusers_hook"):
248
+ module._diffusers_hook.remove_hook(name, recurse=False)
249
+
250
+ def reset_stateful_hooks(self, recurse: bool = True) -> None:
251
+ for hook_name in reversed(self._hook_order):
252
+ hook = self.hooks[hook_name]
253
+ if hook._is_stateful:
254
+ hook.reset_state(self._module_ref)
255
+
256
+ if recurse:
257
+ for module_name, module in unwrap_module(self._module_ref).named_modules():
258
+ if module_name == "":
259
+ continue
260
+ module = unwrap_module(module)
261
+ if hasattr(module, "_diffusers_hook"):
262
+ module._diffusers_hook.reset_stateful_hooks(recurse=False)
263
+
264
+ @classmethod
265
+ def check_if_exists_or_initialize(cls, module: torch.nn.Module) -> "HookRegistry":
266
+ if not hasattr(module, "_diffusers_hook"):
267
+ module._diffusers_hook = cls(module)
268
+ return module._diffusers_hook
269
+
270
+ def _set_context(self, name: str | None = None) -> None:
271
+ for hook_name in reversed(self._hook_order):
272
+ hook = self.hooks[hook_name]
273
+ if hook._is_stateful:
274
+ hook._set_context(self._module_ref, name)
275
+
276
+ for registry in self._get_child_registries():
277
+ registry._set_context(name)
278
+
279
+ def _get_child_registries(self) -> list["HookRegistry"]:
280
+ """Return registries of child modules, using a cached list when available.
281
+
282
+ The cache is built on first call and reused for subsequent calls. This avoids the cost of walking the full
283
+ module tree via named_modules() on every _set_context call, which is significant for large models (e.g. ~2.7ms
284
+ per call on Flux2).
285
+ """
286
+ if not hasattr(self, "_child_registries_cache"):
287
+ self._child_registries_cache = None
288
+
289
+ if self._child_registries_cache is not None:
290
+ return self._child_registries_cache
291
+
292
+ registries = []
293
+ for module_name, module in unwrap_module(self._module_ref).named_modules():
294
+ if module_name == "":
295
+ continue
296
+ module = unwrap_module(module)
297
+ if hasattr(module, "_diffusers_hook"):
298
+ registries.append(module._diffusers_hook)
299
+ self._child_registries_cache = registries
300
+ return registries
301
+
302
+ def __repr__(self) -> str:
303
+ registry_repr = ""
304
+ for i, hook_name in enumerate(self._hook_order):
305
+ if self.hooks[hook_name].__class__.__repr__ is not object.__repr__:
306
+ hook_repr = self.hooks[hook_name].__repr__()
307
+ else:
308
+ hook_repr = self.hooks[hook_name].__class__.__name__
309
+ registry_repr += f" ({i}) {hook_name} - {hook_repr}"
310
+ if i < len(self._hook_order) - 1:
311
+ registry_repr += "\n"
312
+ return f"HookRegistry(\n{registry_repr}\n)"
diffusers/hooks/layer_skip.py ADDED
@@ -0,0 +1,263 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 The HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ import math
16
+ from dataclasses import asdict, dataclass
17
+ from typing import Callable
18
+
19
+ import torch
20
+
21
+ from ..utils import get_logger
22
+ from ..utils.torch_utils import unwrap_module
23
+ from ._common import (
24
+ _ALL_TRANSFORMER_BLOCK_IDENTIFIERS,
25
+ _ATTENTION_CLASSES,
26
+ _FEEDFORWARD_CLASSES,
27
+ _get_submodule_from_fqn,
28
+ )
29
+ from ._helpers import AttentionProcessorRegistry, TransformerBlockRegistry
30
+ from .hooks import HookRegistry, ModelHook
31
+
32
+
33
+ logger = get_logger(__name__) # pylint: disable=invalid-name
34
+
35
+ _LAYER_SKIP_HOOK = "layer_skip_hook"
36
+
37
+
38
+ # Aryan/YiYi TODO: we need to make guider class a config mixin so I think this is not needed
39
+ # either remove or make it serializable
40
+ @dataclass
41
+ class LayerSkipConfig:
42
+ r"""
43
+ Configuration for skipping internal transformer blocks when executing a transformer model.
44
+
45
+ Args:
46
+ indices (`list[int]`):
47
+ The indices of the layer to skip. This is typically the first layer in the transformer block.
48
+ fqn (`str`, defaults to `"auto"`):
49
+ The fully qualified name identifying the stack of transformer blocks. Typically, this is
50
+ `transformer_blocks`, `single_transformer_blocks`, `blocks`, `layers`, or `temporal_transformer_blocks`.
51
+ For automatic detection, set this to `"auto"`. "auto" only works on DiT models. For UNet models, you must
52
+ provide the correct fqn.
53
+ skip_attention (`bool`, defaults to `True`):
54
+ Whether to skip attention blocks.
55
+ skip_ff (`bool`, defaults to `True`):
56
+ Whether to skip feed-forward blocks.
57
+ skip_attention_scores (`bool`, defaults to `False`):
58
+ Whether to skip attention score computation in the attention blocks. This is equivalent to using `value`
59
+ projections as the output of scaled dot product attention.
60
+ dropout (`float`, defaults to `1.0`):
61
+ The dropout probability for dropping the outputs of the skipped layers. By default, this is set to `1.0`,
62
+ meaning that the outputs of the skipped layers are completely ignored. If set to `0.0`, the outputs of the
63
+ skipped layers are fully retained, which is equivalent to not skipping any layers.
64
+ """
65
+
66
+ indices: list[int]
67
+ fqn: str = "auto"
68
+ skip_attention: bool = True
69
+ skip_attention_scores: bool = False
70
+ skip_ff: bool = True
71
+ dropout: float = 1.0
72
+
73
+ def __post_init__(self):
74
+ if not (0 <= self.dropout <= 1):
75
+ raise ValueError(f"Expected `dropout` to be between 0.0 and 1.0, but got {self.dropout}.")
76
+ if not math.isclose(self.dropout, 1.0) and self.skip_attention_scores:
77
+ raise ValueError(
78
+ "Cannot set `skip_attention_scores` to True when `dropout` is not 1.0. Please set `dropout` to 1.0."
79
+ )
80
+
81
+ def to_dict(self):
82
+ return asdict(self)
83
+
84
+ @staticmethod
85
+ def from_dict(data: dict) -> "LayerSkipConfig":
86
+ return LayerSkipConfig(**data)
87
+
88
+
89
+ class AttentionScoreSkipFunctionMode(torch.overrides.TorchFunctionMode):
90
+ def __torch_function__(self, func, types, args=(), kwargs=None):
91
+ if kwargs is None:
92
+ kwargs = {}
93
+ if func is torch.nn.functional.scaled_dot_product_attention:
94
+ query = kwargs.get("query", None)
95
+ key = kwargs.get("key", None)
96
+ value = kwargs.get("value", None)
97
+ query = query if query is not None else args[0]
98
+ key = key if key is not None else args[1]
99
+ value = value if value is not None else args[2]
100
+ # If the Q sequence length does not match KV sequence length, methods like
101
+ # Perturbed Attention Guidance cannot be used (because the caller expects
102
+ # the same sequence length as Q, but if we return V here, it will not match).
103
+ # When Q.shape[2] != V.shape[2], PAG will essentially not be applied and
104
+ # the overall effect would that be of normal CFG with a scale of (guidance_scale + perturbed_guidance_scale).
105
+ if query.shape[2] == value.shape[2]:
106
+ return value
107
+ return func(*args, **kwargs)
108
+
109
+
110
+ class AttentionProcessorSkipHook(ModelHook):
111
+ def __init__(self, skip_processor_output_fn: Callable, skip_attention_scores: bool = False, dropout: float = 1.0):
112
+ self.skip_processor_output_fn = skip_processor_output_fn
113
+ self.skip_attention_scores = skip_attention_scores
114
+ self.dropout = dropout
115
+
116
+ def new_forward(self, module: torch.nn.Module, *args, **kwargs):
117
+ if self.skip_attention_scores:
118
+ if not math.isclose(self.dropout, 1.0):
119
+ raise ValueError(
120
+ "Cannot set `skip_attention_scores` to True when `dropout` is not 1.0. Please set `dropout` to 1.0."
121
+ )
122
+ with AttentionScoreSkipFunctionMode():
123
+ output = self.fn_ref.original_forward(*args, **kwargs)
124
+ else:
125
+ if math.isclose(self.dropout, 1.0):
126
+ output = self.skip_processor_output_fn(module, *args, **kwargs)
127
+ else:
128
+ output = self.fn_ref.original_forward(*args, **kwargs)
129
+ output = torch.nn.functional.dropout(output, p=self.dropout)
130
+ return output
131
+
132
+
133
+ class FeedForwardSkipHook(ModelHook):
134
+ def __init__(self, dropout: float):
135
+ super().__init__()
136
+ self.dropout = dropout
137
+
138
+ def new_forward(self, module: torch.nn.Module, *args, **kwargs):
139
+ if math.isclose(self.dropout, 1.0):
140
+ output = kwargs.get("hidden_states", None)
141
+ if output is None:
142
+ output = kwargs.get("x", None)
143
+ if output is None and len(args) > 0:
144
+ output = args[0]
145
+ else:
146
+ output = self.fn_ref.original_forward(*args, **kwargs)
147
+ output = torch.nn.functional.dropout(output, p=self.dropout)
148
+ return output
149
+
150
+
151
+ class TransformerBlockSkipHook(ModelHook):
152
+ def __init__(self, dropout: float):
153
+ super().__init__()
154
+ self.dropout = dropout
155
+
156
+ def initialize_hook(self, module):
157
+ self._metadata = TransformerBlockRegistry.get(unwrap_module(module).__class__)
158
+ return module
159
+
160
+ def new_forward(self, module: torch.nn.Module, *args, **kwargs):
161
+ if math.isclose(self.dropout, 1.0):
162
+ original_hidden_states = self._metadata._get_parameter_from_args_kwargs("hidden_states", args, kwargs)
163
+ if self._metadata.return_encoder_hidden_states_index is None:
164
+ output = original_hidden_states
165
+ else:
166
+ original_encoder_hidden_states = self._metadata._get_parameter_from_args_kwargs(
167
+ "encoder_hidden_states", args, kwargs
168
+ )
169
+ output = (original_hidden_states, original_encoder_hidden_states)
170
+ else:
171
+ output = self.fn_ref.original_forward(*args, **kwargs)
172
+ output = torch.nn.functional.dropout(output, p=self.dropout)
173
+ return output
174
+
175
+
176
+ def apply_layer_skip(module: torch.nn.Module, config: LayerSkipConfig) -> None:
177
+ r"""
178
+ Apply layer skipping to internal layers of a transformer.
179
+
180
+ Args:
181
+ module (`torch.nn.Module`):
182
+ The transformer model to which the layer skip hook should be applied.
183
+ config (`LayerSkipConfig`):
184
+ The configuration for the layer skip hook.
185
+
186
+ Example:
187
+
188
+ ```python
189
+ >>> from diffusers import apply_layer_skip_hook, CogVideoXTransformer3DModel, LayerSkipConfig
190
+
191
+ >>> transformer = CogVideoXTransformer3DModel.from_pretrained("THUDM/CogVideoX-5b", torch_dtype=torch.bfloat16)
192
+ >>> config = LayerSkipConfig(layer_index=[10, 20], fqn="transformer_blocks")
193
+ >>> apply_layer_skip_hook(transformer, config)
194
+ ```
195
+ """
196
+ _apply_layer_skip_hook(module, config)
197
+
198
+
199
+ def _apply_layer_skip_hook(module: torch.nn.Module, config: LayerSkipConfig, name: str | None = None) -> None:
200
+ name = name or _LAYER_SKIP_HOOK
201
+
202
+ if config.skip_attention and config.skip_attention_scores:
203
+ raise ValueError("Cannot set both `skip_attention` and `skip_attention_scores` to True. Please choose one.")
204
+ if not math.isclose(config.dropout, 1.0) and config.skip_attention_scores:
205
+ raise ValueError(
206
+ "Cannot set `skip_attention_scores` to True when `dropout` is not 1.0. Please set `dropout` to 1.0."
207
+ )
208
+
209
+ if config.fqn == "auto":
210
+ for identifier in _ALL_TRANSFORMER_BLOCK_IDENTIFIERS:
211
+ if hasattr(module, identifier):
212
+ config.fqn = identifier
213
+ break
214
+ else:
215
+ raise ValueError(
216
+ "Could not find a suitable identifier for the transformer blocks automatically. Please provide a valid "
217
+ "`fqn` (fully qualified name) that identifies a stack of transformer blocks."
218
+ )
219
+
220
+ transformer_blocks = _get_submodule_from_fqn(module, config.fqn)
221
+ if transformer_blocks is None or not isinstance(transformer_blocks, torch.nn.ModuleList):
222
+ raise ValueError(
223
+ f"Could not find {config.fqn} in the provided module, or configured `fqn` (fully qualified name) does not identify "
224
+ f"a `torch.nn.ModuleList`. Please provide a valid `fqn` that identifies a stack of transformer blocks."
225
+ )
226
+ if len(config.indices) == 0:
227
+ raise ValueError("Layer index list is empty. Please provide a non-empty list of layer indices to skip.")
228
+
229
+ blocks_found = False
230
+ for i, block in enumerate(transformer_blocks):
231
+ if i not in config.indices:
232
+ continue
233
+
234
+ blocks_found = True
235
+
236
+ if config.skip_attention and config.skip_ff:
237
+ logger.debug(f"Applying TransformerBlockSkipHook to '{config.fqn}.{i}'")
238
+ registry = HookRegistry.check_if_exists_or_initialize(block)
239
+ hook = TransformerBlockSkipHook(config.dropout)
240
+ registry.register_hook(hook, name)
241
+
242
+ elif config.skip_attention or config.skip_attention_scores:
243
+ for submodule_name, submodule in block.named_modules():
244
+ if isinstance(submodule, _ATTENTION_CLASSES) and not submodule.is_cross_attention:
245
+ logger.debug(f"Applying AttentionProcessorSkipHook to '{config.fqn}.{i}.{submodule_name}'")
246
+ output_fn = AttentionProcessorRegistry.get(submodule.processor.__class__).skip_processor_output_fn
247
+ registry = HookRegistry.check_if_exists_or_initialize(submodule)
248
+ hook = AttentionProcessorSkipHook(output_fn, config.skip_attention_scores, config.dropout)
249
+ registry.register_hook(hook, name)
250
+
251
+ if config.skip_ff:
252
+ for submodule_name, submodule in block.named_modules():
253
+ if isinstance(submodule, _FEEDFORWARD_CLASSES):
254
+ logger.debug(f"Applying FeedForwardSkipHook to '{config.fqn}.{i}.{submodule_name}'")
255
+ registry = HookRegistry.check_if_exists_or_initialize(submodule)
256
+ hook = FeedForwardSkipHook(config.dropout)
257
+ registry.register_hook(hook, name)
258
+
259
+ if not blocks_found:
260
+ raise ValueError(
261
+ f"Could not find any transformer blocks matching the provided indices {config.indices} and "
262
+ f"fully qualified name '{config.fqn}'. Please check the indices and fqn for correctness."
263
+ )
diffusers/hooks/layerwise_casting.py ADDED
@@ -0,0 +1,240 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 The HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ import re
16
+ from typing import Type
17
+
18
+ import torch
19
+
20
+ from ..utils import get_logger, is_peft_available, is_peft_version
21
+ from ._common import _GO_LC_SUPPORTED_PYTORCH_LAYERS
22
+ from .hooks import HookRegistry, ModelHook
23
+
24
+
25
+ logger = get_logger(__name__) # pylint: disable=invalid-name
26
+
27
+
28
+ # fmt: off
29
+ _LAYERWISE_CASTING_HOOK = "layerwise_casting"
30
+ _PEFT_AUTOCAST_DISABLE_HOOK = "peft_autocast_disable"
31
+ DEFAULT_SKIP_MODULES_PATTERN = ("pos_embed", "patch_embed", "norm", "^proj_in$", "^proj_out$")
32
+ # fmt: on
33
+
34
+ _SHOULD_DISABLE_PEFT_INPUT_AUTOCAST = is_peft_available() and is_peft_version(">", "0.14.0")
35
+ if _SHOULD_DISABLE_PEFT_INPUT_AUTOCAST:
36
+ from peft.helpers import disable_input_dtype_casting
37
+ from peft.tuners.tuners_utils import BaseTunerLayer
38
+
39
+
40
+ class LayerwiseCastingHook(ModelHook):
41
+ r"""
42
+ A hook that casts the weights of a module to a high precision dtype for computation, and to a low precision dtype
43
+ for storage. This process may lead to quality loss in the output, but can significantly reduce the memory
44
+ footprint.
45
+ """
46
+
47
+ _is_stateful = False
48
+
49
+ def __init__(self, storage_dtype: torch.dtype, compute_dtype: torch.dtype, non_blocking: bool) -> None:
50
+ self.storage_dtype = storage_dtype
51
+ self.compute_dtype = compute_dtype
52
+ self.non_blocking = non_blocking
53
+
54
+ def initialize_hook(self, module: torch.nn.Module):
55
+ module.to(dtype=self.storage_dtype, non_blocking=self.non_blocking)
56
+ return module
57
+
58
+ def deinitalize_hook(self, module: torch.nn.Module):
59
+ raise NotImplementedError(
60
+ "LayerwiseCastingHook does not support deinitialization. A model once enabled with layerwise casting will "
61
+ "have casted its weights to a lower precision dtype for storage. Casting this back to the original dtype "
62
+ "will lead to precision loss, which might have an impact on the model's generation quality. The model should "
63
+ "be re-initialized and loaded in the original dtype."
64
+ )
65
+
66
+ def pre_forward(self, module: torch.nn.Module, *args, **kwargs):
67
+ module.to(dtype=self.compute_dtype, non_blocking=self.non_blocking)
68
+ return args, kwargs
69
+
70
+ def post_forward(self, module: torch.nn.Module, output):
71
+ module.to(dtype=self.storage_dtype, non_blocking=self.non_blocking)
72
+ return output
73
+
74
+
75
+ class PeftInputAutocastDisableHook(ModelHook):
76
+ r"""
77
+ A hook that disables the casting of inputs to the module weight dtype during the forward pass. By default, PEFT
78
+ casts the inputs to the weight dtype of the module, which can lead to precision loss.
79
+
80
+ The reasons for needing this are:
81
+ - If we don't add PEFT layers' weight names to `skip_modules_pattern` when applying layerwise casting, the
82
+ inputs will be casted to the, possibly lower precision, storage dtype. Reference:
83
+ https://github.com/huggingface/peft/blob/0facdebf6208139cbd8f3586875acb378813dd97/src/peft/tuners/lora/layer.py#L706
84
+ - We can, on our end, use something like accelerate's `send_to_device` but for dtypes. This way, we can ensure
85
+ that the inputs are casted to the computation dtype correctly always. However, there are two goals we are
86
+ hoping to achieve:
87
+ 1. Making forward implementations independent of device/dtype casting operations as much as possible.
88
+ 2. Performing inference without losing information from casting to different precisions. With the current
89
+ PEFT implementation (as linked in the reference above), and assuming running layerwise casting inference
90
+ with storage_dtype=torch.float8_e4m3fn and compute_dtype=torch.bfloat16, inputs are cast to
91
+ torch.float8_e4m3fn in the lora layer. We will then upcast back to torch.bfloat16 when we continue the
92
+ forward pass in PEFT linear forward or Diffusers layer forward, with a `send_to_dtype` operation from
93
+ LayerwiseCastingHook. This will be a lossy operation and result in poorer generation quality.
94
+ """
95
+
96
+ def new_forward(self, module: torch.nn.Module, *args, **kwargs):
97
+ with disable_input_dtype_casting(module):
98
+ return self.fn_ref.original_forward(*args, **kwargs)
99
+
100
+
101
+ def apply_layerwise_casting(
102
+ module: torch.nn.Module,
103
+ storage_dtype: torch.dtype,
104
+ compute_dtype: torch.dtype,
105
+ skip_modules_pattern: str | tuple[str, ...] = "auto",
106
+ skip_modules_classes: tuple[Type[torch.nn.Module], ...] | None = None,
107
+ non_blocking: bool = False,
108
+ ) -> None:
109
+ r"""
110
+ Applies layerwise casting to a given module. The module expected here is a Diffusers ModelMixin but it can be any
111
+ nn.Module using diffusers layers or pytorch primitives.
112
+
113
+ Example:
114
+
115
+ ```python
116
+ >>> import torch
117
+ >>> from diffusers import CogVideoXTransformer3DModel
118
+
119
+ >>> transformer = CogVideoXTransformer3DModel.from_pretrained(
120
+ ... model_id, subfolder="transformer", torch_dtype=torch.bfloat16
121
+ ... )
122
+
123
+ >>> apply_layerwise_casting(
124
+ ... transformer,
125
+ ... storage_dtype=torch.float8_e4m3fn,
126
+ ... compute_dtype=torch.bfloat16,
127
+ ... skip_modules_pattern=["patch_embed", "norm", "proj_out"],
128
+ ... non_blocking=True,
129
+ ... )
130
+ ```
131
+
132
+ Args:
133
+ module (`torch.nn.Module`):
134
+ The module whose leaf modules will be cast to a high precision dtype for computation, and to a low
135
+ precision dtype for storage.
136
+ storage_dtype (`torch.dtype`):
137
+ The dtype to cast the module to before/after the forward pass for storage.
138
+ compute_dtype (`torch.dtype`):
139
+ The dtype to cast the module to during the forward pass for computation.
140
+ skip_modules_pattern (`tuple[str, ...]`, defaults to `"auto"`):
141
+ A list of patterns to match the names of the modules to skip during the layerwise casting process. If set
142
+ to `"auto"`, the default patterns are used. If set to `None`, no modules are skipped. If set to `None`
143
+ alongside `skip_modules_classes` being `None`, the layerwise casting is applied directly to the module
144
+ instead of its internal submodules.
145
+ skip_modules_classes (`tuple[Type[torch.nn.Module], ...]`, defaults to `None`):
146
+ A list of module classes to skip during the layerwise casting process.
147
+ non_blocking (`bool`, defaults to `False`):
148
+ If `True`, the weight casting operations are non-blocking.
149
+ """
150
+ if skip_modules_pattern == "auto":
151
+ skip_modules_pattern = DEFAULT_SKIP_MODULES_PATTERN
152
+
153
+ if skip_modules_classes is None and skip_modules_pattern is None:
154
+ apply_layerwise_casting_hook(module, storage_dtype, compute_dtype, non_blocking)
155
+ return
156
+
157
+ _apply_layerwise_casting(
158
+ module,
159
+ storage_dtype,
160
+ compute_dtype,
161
+ skip_modules_pattern,
162
+ skip_modules_classes,
163
+ non_blocking,
164
+ )
165
+ _disable_peft_input_autocast(module)
166
+
167
+
168
+ def _apply_layerwise_casting(
169
+ module: torch.nn.Module,
170
+ storage_dtype: torch.dtype,
171
+ compute_dtype: torch.dtype,
172
+ skip_modules_pattern: tuple[str, ...] | None = None,
173
+ skip_modules_classes: tuple[Type[torch.nn.Module], ...] | None = None,
174
+ non_blocking: bool = False,
175
+ _prefix: str = "",
176
+ ) -> None:
177
+ should_skip = (skip_modules_classes is not None and isinstance(module, skip_modules_classes)) or (
178
+ skip_modules_pattern is not None and any(re.search(pattern, _prefix) for pattern in skip_modules_pattern)
179
+ )
180
+ if should_skip:
181
+ logger.debug(f'Skipping layerwise casting for layer "{_prefix}"')
182
+ return
183
+
184
+ if isinstance(module, _GO_LC_SUPPORTED_PYTORCH_LAYERS):
185
+ logger.debug(f'Applying layerwise casting to layer "{_prefix}"')
186
+ apply_layerwise_casting_hook(module, storage_dtype, compute_dtype, non_blocking)
187
+ return
188
+
189
+ for name, submodule in module.named_children():
190
+ layer_name = f"{_prefix}.{name}" if _prefix else name
191
+ _apply_layerwise_casting(
192
+ submodule,
193
+ storage_dtype,
194
+ compute_dtype,
195
+ skip_modules_pattern,
196
+ skip_modules_classes,
197
+ non_blocking,
198
+ _prefix=layer_name,
199
+ )
200
+
201
+
202
+ def apply_layerwise_casting_hook(
203
+ module: torch.nn.Module, storage_dtype: torch.dtype, compute_dtype: torch.dtype, non_blocking: bool
204
+ ) -> None:
205
+ r"""
206
+ Applies a `LayerwiseCastingHook` to a given module.
207
+
208
+ Args:
209
+ module (`torch.nn.Module`):
210
+ The module to attach the hook to.
211
+ storage_dtype (`torch.dtype`):
212
+ The dtype to cast the module to before the forward pass.
213
+ compute_dtype (`torch.dtype`):
214
+ The dtype to cast the module to during the forward pass.
215
+ non_blocking (`bool`):
216
+ If `True`, the weight casting operations are non-blocking.
217
+ """
218
+ registry = HookRegistry.check_if_exists_or_initialize(module)
219
+ hook = LayerwiseCastingHook(storage_dtype, compute_dtype, non_blocking)
220
+ registry.register_hook(hook, _LAYERWISE_CASTING_HOOK)
221
+
222
+
223
+ def _is_layerwise_casting_active(module: torch.nn.Module) -> bool:
224
+ for submodule in module.modules():
225
+ if (
226
+ hasattr(submodule, "_diffusers_hook")
227
+ and submodule._diffusers_hook.get_hook(_LAYERWISE_CASTING_HOOK) is not None
228
+ ):
229
+ return True
230
+ return False
231
+
232
+
233
+ def _disable_peft_input_autocast(module: torch.nn.Module) -> None:
234
+ if not _SHOULD_DISABLE_PEFT_INPUT_AUTOCAST:
235
+ return
236
+ for submodule in module.modules():
237
+ if isinstance(submodule, BaseTunerLayer) and _is_layerwise_casting_active(submodule):
238
+ registry = HookRegistry.check_if_exists_or_initialize(submodule)
239
+ hook = PeftInputAutocastDisableHook()
240
+ registry.register_hook(hook, _PEFT_AUTOCAST_DISABLE_HOOK)
diffusers/hooks/mag_cache.py ADDED
@@ -0,0 +1,468 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 The HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ from dataclasses import dataclass
16
+ from typing import List, Optional, Tuple, Union
17
+
18
+ import torch
19
+
20
+ from ..utils import get_logger
21
+ from ..utils.torch_utils import unwrap_module
22
+ from ._common import _ALL_TRANSFORMER_BLOCK_IDENTIFIERS
23
+ from ._helpers import TransformerBlockRegistry
24
+ from .hooks import BaseState, HookRegistry, ModelHook, StateManager
25
+
26
+
27
+ logger = get_logger(__name__) # pylint: disable=invalid-name
28
+
29
+ _MAG_CACHE_LEADER_BLOCK_HOOK = "mag_cache_leader_block_hook"
30
+ _MAG_CACHE_BLOCK_HOOK = "mag_cache_block_hook"
31
+
32
+ # Default Mag Ratios for Flux models (Dev/Schnell) are provided for convenience.
33
+ # Users must explicitly pass these to the config if using Flux.
34
+ # Reference: https://github.com/Zehong-Ma/MagCache
35
+ FLUX_MAG_RATIOS = torch.tensor(
36
+ [1.0]
37
+ + [
38
+ 1.21094,
39
+ 1.11719,
40
+ 1.07812,
41
+ 1.0625,
42
+ 1.03906,
43
+ 1.03125,
44
+ 1.03906,
45
+ 1.02344,
46
+ 1.03125,
47
+ 1.02344,
48
+ 0.98047,
49
+ 1.01562,
50
+ 1.00781,
51
+ 1.0,
52
+ 1.00781,
53
+ 1.0,
54
+ 1.00781,
55
+ 1.0,
56
+ 1.0,
57
+ 0.99609,
58
+ 0.99609,
59
+ 0.98047,
60
+ 0.98828,
61
+ 0.96484,
62
+ 0.95703,
63
+ 0.93359,
64
+ 0.89062,
65
+ ]
66
+ )
67
+
68
+
69
+ def nearest_interp(src_array: torch.Tensor, target_length: int) -> torch.Tensor:
70
+ """
71
+ Interpolate the source array to the target length using nearest neighbor interpolation.
72
+ """
73
+ src_length = len(src_array)
74
+ if target_length == 1:
75
+ return src_array[-1:]
76
+
77
+ scale = (src_length - 1) / (target_length - 1)
78
+ grid = torch.arange(target_length, device=src_array.device, dtype=torch.float32)
79
+ mapped_indices = torch.round(grid * scale).long()
80
+ return src_array[mapped_indices]
81
+
82
+
83
+ @dataclass
84
+ class MagCacheConfig:
85
+ r"""
86
+ Configuration for [MagCache](https://github.com/Zehong-Ma/MagCache).
87
+
88
+ Args:
89
+ threshold (`float`, defaults to `0.06`):
90
+ The threshold for the accumulated error. If the accumulated error is below this threshold, the block
91
+ computation is skipped. A higher threshold allows for more aggressive skipping (faster) but may degrade
92
+ quality.
93
+ max_skip_steps (`int`, defaults to `3`):
94
+ The maximum number of consecutive steps that can be skipped (K in the paper).
95
+ retention_ratio (`float`, defaults to `0.2`):
96
+ The fraction of initial steps during which skipping is disabled to ensure stability. For example, if
97
+ `num_inference_steps` is 28 and `retention_ratio` is 0.2, the first 6 steps will never be skipped.
98
+ num_inference_steps (`int`, defaults to `28`):
99
+ The number of inference steps used in the pipeline. This is required to interpolate `mag_ratios` correctly.
100
+ mag_ratios (`torch.Tensor`, *optional*):
101
+ The pre-computed magnitude ratios for the model. These are checkpoint-dependent. If not provided, you must
102
+ set `calibrate=True` to calculate them for your specific model. For Flux models, you can use
103
+ `diffusers.hooks.mag_cache.FLUX_MAG_RATIOS`.
104
+ calibrate (`bool`, defaults to `False`):
105
+ If True, enables calibration mode. In this mode, no blocks are skipped. Instead, the hook calculates the
106
+ magnitude ratios for the current run and logs them at the end. Use this to obtain `mag_ratios` for new
107
+ models or schedulers.
108
+ """
109
+
110
+ threshold: float = 0.06
111
+ max_skip_steps: int = 3
112
+ retention_ratio: float = 0.2
113
+ num_inference_steps: int = 28
114
+ mag_ratios: Optional[Union[torch.Tensor, List[float]]] = None
115
+ calibrate: bool = False
116
+
117
+ def __post_init__(self):
118
+ # User MUST provide ratios OR enable calibration.
119
+ if self.mag_ratios is None and not self.calibrate:
120
+ raise ValueError(
121
+ " `mag_ratios` must be provided for MagCache inference because these ratios are model-dependent.\n"
122
+ "To get them for your model:\n"
123
+ "1. Initialize `MagCacheConfig(calibrate=True, ...)`\n"
124
+ "2. Run inference on your model once.\n"
125
+ "3. Copy the printed ratios array and pass it to `mag_ratios` in the config.\n"
126
+ "For Flux models, you can import `FLUX_MAG_RATIOS` from `diffusers.hooks.mag_cache`."
127
+ )
128
+
129
+ if not self.calibrate and self.mag_ratios is not None:
130
+ if not torch.is_tensor(self.mag_ratios):
131
+ self.mag_ratios = torch.tensor(self.mag_ratios)
132
+
133
+ if len(self.mag_ratios) != self.num_inference_steps:
134
+ logger.debug(
135
+ f"Interpolating mag_ratios from length {len(self.mag_ratios)} to {self.num_inference_steps}"
136
+ )
137
+ self.mag_ratios = nearest_interp(self.mag_ratios, self.num_inference_steps)
138
+
139
+
140
+ class MagCacheState(BaseState):
141
+ def __init__(self) -> None:
142
+ super().__init__()
143
+ # Cache for the residual (output - input) from the *previous* timestep
144
+ self.previous_residual: torch.Tensor = None
145
+
146
+ # State inputs/outputs for the current forward pass
147
+ self.head_block_input: Union[torch.Tensor, Tuple[torch.Tensor, ...]] = None
148
+ self.should_compute: bool = True
149
+
150
+ # MagCache accumulators
151
+ self.accumulated_ratio: float = 1.0
152
+ self.accumulated_err: float = 0.0
153
+ self.accumulated_steps: int = 0
154
+
155
+ # Current step counter (timestep index)
156
+ self.step_index: int = 0
157
+
158
+ # Calibration storage
159
+ self.calibration_ratios: List[float] = []
160
+
161
+ def reset(self):
162
+ self.previous_residual = None
163
+ self.should_compute = True
164
+ self.accumulated_ratio = 1.0
165
+ self.accumulated_err = 0.0
166
+ self.accumulated_steps = 0
167
+ self.step_index = 0
168
+ self.calibration_ratios = []
169
+
170
+
171
+ class MagCacheHeadHook(ModelHook):
172
+ _is_stateful = True
173
+
174
+ def __init__(self, state_manager: StateManager, config: MagCacheConfig):
175
+ self.state_manager = state_manager
176
+ self.config = config
177
+ self._metadata = None
178
+
179
+ def initialize_hook(self, module):
180
+ unwrapped_module = unwrap_module(module)
181
+ self._metadata = TransformerBlockRegistry.get(unwrapped_module.__class__)
182
+ return module
183
+
184
+ @torch.compiler.disable
185
+ def new_forward(self, module: torch.nn.Module, *args, **kwargs):
186
+ if self.state_manager._current_context is None:
187
+ self.state_manager.set_context("inference")
188
+
189
+ arg_name = self._metadata.hidden_states_argument_name
190
+ hidden_states = self._metadata._get_parameter_from_args_kwargs(arg_name, args, kwargs)
191
+
192
+ state: MagCacheState = self.state_manager.get_state()
193
+ state.head_block_input = hidden_states
194
+
195
+ should_compute = True
196
+
197
+ if self.config.calibrate:
198
+ # Never skip during calibration
199
+ should_compute = True
200
+ else:
201
+ # MagCache Logic
202
+ current_step = state.step_index
203
+ if current_step >= len(self.config.mag_ratios):
204
+ current_scale = 1.0
205
+ else:
206
+ current_scale = self.config.mag_ratios[current_step]
207
+
208
+ retention_step = int(self.config.retention_ratio * self.config.num_inference_steps + 0.5)
209
+
210
+ if current_step >= retention_step:
211
+ state.accumulated_ratio *= current_scale
212
+ state.accumulated_steps += 1
213
+ state.accumulated_err += abs(1.0 - state.accumulated_ratio)
214
+
215
+ if (
216
+ state.previous_residual is not None
217
+ and state.accumulated_err <= self.config.threshold
218
+ and state.accumulated_steps <= self.config.max_skip_steps
219
+ ):
220
+ should_compute = False
221
+ else:
222
+ state.accumulated_ratio = 1.0
223
+ state.accumulated_steps = 0
224
+ state.accumulated_err = 0.0
225
+
226
+ state.should_compute = should_compute
227
+
228
+ if not should_compute:
229
+ logger.debug(f"MagCache: Skipping step {state.step_index}")
230
+ # Apply MagCache: Output = Input + Previous Residual
231
+
232
+ output = hidden_states
233
+ res = state.previous_residual
234
+
235
+ if res.device != output.device:
236
+ res = res.to(output.device)
237
+
238
+ # Attempt to apply residual handling shape mismatches (e.g., text+image vs image only)
239
+ if res.shape == output.shape:
240
+ output = output + res
241
+ elif (
242
+ output.ndim == 3
243
+ and res.ndim == 3
244
+ and output.shape[0] == res.shape[0]
245
+ and output.shape[2] == res.shape[2]
246
+ ):
247
+ # Assuming concatenation where image part is at the end (standard in Flux/SD3)
248
+ diff = output.shape[1] - res.shape[1]
249
+ if diff > 0:
250
+ output = output.clone()
251
+ output[:, diff:, :] = output[:, diff:, :] + res
252
+ else:
253
+ logger.warning(
254
+ f"MagCache: Dimension mismatch. Input {output.shape}, Residual {res.shape}. "
255
+ "Cannot apply residual safely. Returning input without residual."
256
+ )
257
+ else:
258
+ logger.warning(
259
+ f"MagCache: Dimension mismatch. Input {output.shape}, Residual {res.shape}. "
260
+ "Cannot apply residual safely. Returning input without residual."
261
+ )
262
+
263
+ if self._metadata.return_encoder_hidden_states_index is not None:
264
+ original_encoder_hidden_states = self._metadata._get_parameter_from_args_kwargs(
265
+ "encoder_hidden_states", args, kwargs
266
+ )
267
+ max_idx = max(
268
+ self._metadata.return_hidden_states_index, self._metadata.return_encoder_hidden_states_index
269
+ )
270
+ ret_list = [None] * (max_idx + 1)
271
+ ret_list[self._metadata.return_hidden_states_index] = output
272
+ ret_list[self._metadata.return_encoder_hidden_states_index] = original_encoder_hidden_states
273
+ return tuple(ret_list)
274
+ else:
275
+ return output
276
+
277
+ else:
278
+ # Compute original forward
279
+ output = self.fn_ref.original_forward(*args, **kwargs)
280
+ return output
281
+
282
+ def reset_state(self, module):
283
+ self.state_manager.reset()
284
+ return module
285
+
286
+
287
+ class MagCacheBlockHook(ModelHook):
288
+ def __init__(self, state_manager: StateManager, is_tail: bool = False, config: MagCacheConfig = None):
289
+ super().__init__()
290
+ self.state_manager = state_manager
291
+ self.is_tail = is_tail
292
+ self.config = config
293
+ self._metadata = None
294
+
295
+ def initialize_hook(self, module):
296
+ unwrapped_module = unwrap_module(module)
297
+ self._metadata = TransformerBlockRegistry.get(unwrapped_module.__class__)
298
+ return module
299
+
300
+ @torch.compiler.disable
301
+ def new_forward(self, module: torch.nn.Module, *args, **kwargs):
302
+ if self.state_manager._current_context is None:
303
+ self.state_manager.set_context("inference")
304
+ state: MagCacheState = self.state_manager.get_state()
305
+
306
+ if not state.should_compute:
307
+ arg_name = self._metadata.hidden_states_argument_name
308
+ hidden_states = self._metadata._get_parameter_from_args_kwargs(arg_name, args, kwargs)
309
+
310
+ if self.is_tail:
311
+ # Still need to advance step index even if we skip
312
+ self._advance_step(state)
313
+
314
+ if self._metadata.return_encoder_hidden_states_index is not None:
315
+ encoder_hidden_states = self._metadata._get_parameter_from_args_kwargs(
316
+ "encoder_hidden_states", args, kwargs
317
+ )
318
+ max_idx = max(
319
+ self._metadata.return_hidden_states_index, self._metadata.return_encoder_hidden_states_index
320
+ )
321
+ ret_list = [None] * (max_idx + 1)
322
+ ret_list[self._metadata.return_hidden_states_index] = hidden_states
323
+ ret_list[self._metadata.return_encoder_hidden_states_index] = encoder_hidden_states
324
+ return tuple(ret_list)
325
+
326
+ return hidden_states
327
+
328
+ output = self.fn_ref.original_forward(*args, **kwargs)
329
+
330
+ if self.is_tail:
331
+ # Calculate residual for next steps
332
+ if isinstance(output, tuple):
333
+ out_hidden = output[self._metadata.return_hidden_states_index]
334
+ else:
335
+ out_hidden = output
336
+
337
+ in_hidden = state.head_block_input
338
+
339
+ if in_hidden is None:
340
+ return output
341
+
342
+ # Determine residual
343
+ if out_hidden.shape == in_hidden.shape:
344
+ residual = out_hidden - in_hidden
345
+ elif out_hidden.ndim == 3 and in_hidden.ndim == 3 and out_hidden.shape[2] == in_hidden.shape[2]:
346
+ diff = in_hidden.shape[1] - out_hidden.shape[1]
347
+ if diff == 0:
348
+ residual = out_hidden - in_hidden
349
+ else:
350
+ residual = out_hidden - in_hidden # Fallback to matching tail
351
+ else:
352
+ # Fallback for completely mismatched shapes
353
+ residual = out_hidden
354
+
355
+ if self.config.calibrate:
356
+ self._perform_calibration_step(state, residual)
357
+
358
+ state.previous_residual = residual
359
+ self._advance_step(state)
360
+
361
+ return output
362
+
363
+ def _perform_calibration_step(self, state: MagCacheState, current_residual: torch.Tensor):
364
+ if state.previous_residual is None:
365
+ # First step has no previous residual to compare against.
366
+ # log 1.0 as a neutral starting point.
367
+ ratio = 1.0
368
+ else:
369
+ # MagCache Calibration Formula: mean(norm(curr) / norm(prev))
370
+ # norm(dim=-1) gives magnitude of each token vector
371
+ curr_norm = torch.linalg.norm(current_residual.float(), dim=-1)
372
+ prev_norm = torch.linalg.norm(state.previous_residual.float(), dim=-1)
373
+
374
+ # Avoid division by zero
375
+ ratio = (curr_norm / (prev_norm + 1e-8)).mean().item()
376
+
377
+ state.calibration_ratios.append(ratio)
378
+
379
+ def _advance_step(self, state: MagCacheState):
380
+ state.step_index += 1
381
+ if state.step_index >= self.config.num_inference_steps:
382
+ # End of inference loop
383
+ if self.config.calibrate:
384
+ print("\n[MagCache] Calibration Complete. Copy these values to MagCacheConfig(mag_ratios=...):")
385
+ print(f"{state.calibration_ratios}\n")
386
+ logger.info(f"MagCache Calibration Results: {state.calibration_ratios}")
387
+
388
+ # Reset state
389
+ state.step_index = 0
390
+ state.accumulated_ratio = 1.0
391
+ state.accumulated_steps = 0
392
+ state.accumulated_err = 0.0
393
+ state.previous_residual = None
394
+ state.calibration_ratios = []
395
+
396
+
397
+ def apply_mag_cache(module: torch.nn.Module, config: MagCacheConfig) -> None:
398
+ """
399
+ Applies MagCache to a given module (typically a Transformer).
400
+
401
+ Args:
402
+ module (`torch.nn.Module`):
403
+ The module to apply MagCache to.
404
+ config (`MagCacheConfig`):
405
+ The configuration for MagCache.
406
+ """
407
+ # Initialize registry on the root module so the Pipeline can set context.
408
+ HookRegistry.check_if_exists_or_initialize(module)
409
+
410
+ state_manager = StateManager(MagCacheState, (), {})
411
+ remaining_blocks = []
412
+
413
+ for name, submodule in module.named_children():
414
+ if name not in _ALL_TRANSFORMER_BLOCK_IDENTIFIERS or not isinstance(submodule, torch.nn.ModuleList):
415
+ continue
416
+ for index, block in enumerate(submodule):
417
+ remaining_blocks.append((f"{name}.{index}", block))
418
+
419
+ if not remaining_blocks:
420
+ logger.warning("MagCache: No transformer blocks found to apply hooks.")
421
+ return
422
+
423
+ # Handle single-block models
424
+ if len(remaining_blocks) == 1:
425
+ name, block = remaining_blocks[0]
426
+ logger.info(f"MagCache: Applying Head+Tail Hooks to single block '{name}'")
427
+ _apply_mag_cache_block_hook(block, state_manager, config, is_tail=True)
428
+ _apply_mag_cache_head_hook(block, state_manager, config)
429
+ return
430
+
431
+ head_block_name, head_block = remaining_blocks.pop(0)
432
+ tail_block_name, tail_block = remaining_blocks.pop(-1)
433
+
434
+ logger.info(f"MagCache: Applying Head Hook to {head_block_name}")
435
+ _apply_mag_cache_head_hook(head_block, state_manager, config)
436
+
437
+ for name, block in remaining_blocks:
438
+ _apply_mag_cache_block_hook(block, state_manager, config)
439
+
440
+ logger.info(f"MagCache: Applying Tail Hook to {tail_block_name}")
441
+ _apply_mag_cache_block_hook(tail_block, state_manager, config, is_tail=True)
442
+
443
+
444
+ def _apply_mag_cache_head_hook(block: torch.nn.Module, state_manager: StateManager, config: MagCacheConfig) -> None:
445
+ registry = HookRegistry.check_if_exists_or_initialize(block)
446
+
447
+ # Automatically remove existing hook to allow re-application (e.g. switching modes)
448
+ if registry.get_hook(_MAG_CACHE_LEADER_BLOCK_HOOK) is not None:
449
+ registry.remove_hook(_MAG_CACHE_LEADER_BLOCK_HOOK)
450
+
451
+ hook = MagCacheHeadHook(state_manager, config)
452
+ registry.register_hook(hook, _MAG_CACHE_LEADER_BLOCK_HOOK)
453
+
454
+
455
+ def _apply_mag_cache_block_hook(
456
+ block: torch.nn.Module,
457
+ state_manager: StateManager,
458
+ config: MagCacheConfig,
459
+ is_tail: bool = False,
460
+ ) -> None:
461
+ registry = HookRegistry.check_if_exists_or_initialize(block)
462
+
463
+ # Automatically remove existing hook to allow re-application
464
+ if registry.get_hook(_MAG_CACHE_BLOCK_HOOK) is not None:
465
+ registry.remove_hook(_MAG_CACHE_BLOCK_HOOK)
466
+
467
+ hook = MagCacheBlockHook(state_manager, is_tail, config)
468
+ registry.register_hook(hook, _MAG_CACHE_BLOCK_HOOK)
diffusers/hooks/pyramid_attention_broadcast.py ADDED
@@ -0,0 +1,314 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 The HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ import re
16
+ from dataclasses import dataclass
17
+ from typing import Any, Callable
18
+
19
+ import torch
20
+
21
+ from ..models.attention import AttentionModuleMixin
22
+ from ..models.attention_processor import Attention, MochiAttention
23
+ from ..utils import logging
24
+ from ._common import (
25
+ _ATTENTION_CLASSES,
26
+ _CROSS_TRANSFORMER_BLOCK_IDENTIFIERS,
27
+ _SPATIAL_TRANSFORMER_BLOCK_IDENTIFIERS,
28
+ _TEMPORAL_TRANSFORMER_BLOCK_IDENTIFIERS,
29
+ )
30
+ from .hooks import HookRegistry, ModelHook
31
+
32
+
33
+ logger = logging.get_logger(__name__) # pylint: disable=invalid-name
34
+
35
+
36
+ _PYRAMID_ATTENTION_BROADCAST_HOOK = "pyramid_attention_broadcast"
37
+
38
+
39
+ @dataclass
40
+ class PyramidAttentionBroadcastConfig:
41
+ r"""
42
+ Configuration for Pyramid Attention Broadcast.
43
+
44
+ Args:
45
+ spatial_attention_block_skip_range (`int`, *optional*, defaults to `None`):
46
+ The number of times a specific spatial attention broadcast is skipped before computing the attention states
47
+ to re-use. If this is set to the value `N`, the attention computation will be skipped `N - 1` times (i.e.,
48
+ old attention states will be reused) before computing the new attention states again.
49
+ temporal_attention_block_skip_range (`int`, *optional*, defaults to `None`):
50
+ The number of times a specific temporal attention broadcast is skipped before computing the attention
51
+ states to re-use. If this is set to the value `N`, the attention computation will be skipped `N - 1` times
52
+ (i.e., old attention states will be reused) before computing the new attention states again.
53
+ cross_attention_block_skip_range (`int`, *optional*, defaults to `None`):
54
+ The number of times a specific cross-attention broadcast is skipped before computing the attention states
55
+ to re-use. If this is set to the value `N`, the attention computation will be skipped `N - 1` times (i.e.,
56
+ old attention states will be reused) before computing the new attention states again.
57
+ spatial_attention_timestep_skip_range (`tuple[int, int]`, defaults to `(100, 800)`):
58
+ The range of timesteps to skip in the spatial attention layer. The attention computations will be
59
+ conditionally skipped if the current timestep is within the specified range.
60
+ temporal_attention_timestep_skip_range (`tuple[int, int]`, defaults to `(100, 800)`):
61
+ The range of timesteps to skip in the temporal attention layer. The attention computations will be
62
+ conditionally skipped if the current timestep is within the specified range.
63
+ cross_attention_timestep_skip_range (`tuple[int, int]`, defaults to `(100, 800)`):
64
+ The range of timesteps to skip in the cross-attention layer. The attention computations will be
65
+ conditionally skipped if the current timestep is within the specified range.
66
+ spatial_attention_block_identifiers (`tuple[str, ...]`):
67
+ The identifiers to match against the layer names to determine if the layer is a spatial attention layer.
68
+ temporal_attention_block_identifiers (`tuple[str, ...]`):
69
+ The identifiers to match against the layer names to determine if the layer is a temporal attention layer.
70
+ cross_attention_block_identifiers (`tuple[str, ...]`):
71
+ The identifiers to match against the layer names to determine if the layer is a cross-attention layer.
72
+ """
73
+
74
+ spatial_attention_block_skip_range: int | None = None
75
+ temporal_attention_block_skip_range: int | None = None
76
+ cross_attention_block_skip_range: int | None = None
77
+
78
+ spatial_attention_timestep_skip_range: tuple[int, int] = (100, 800)
79
+ temporal_attention_timestep_skip_range: tuple[int, int] = (100, 800)
80
+ cross_attention_timestep_skip_range: tuple[int, int] = (100, 800)
81
+
82
+ spatial_attention_block_identifiers: tuple[str, ...] = _SPATIAL_TRANSFORMER_BLOCK_IDENTIFIERS
83
+ temporal_attention_block_identifiers: tuple[str, ...] = _TEMPORAL_TRANSFORMER_BLOCK_IDENTIFIERS
84
+ cross_attention_block_identifiers: tuple[str, ...] = _CROSS_TRANSFORMER_BLOCK_IDENTIFIERS
85
+
86
+ current_timestep_callback: Callable[[], int] = None
87
+
88
+ # TODO(aryan): add PAB for MLP layers (very limited speedup from testing with original codebase
89
+ # so not added for now)
90
+
91
+ def __repr__(self) -> str:
92
+ return (
93
+ f"PyramidAttentionBroadcastConfig(\n"
94
+ f" spatial_attention_block_skip_range={self.spatial_attention_block_skip_range},\n"
95
+ f" temporal_attention_block_skip_range={self.temporal_attention_block_skip_range},\n"
96
+ f" cross_attention_block_skip_range={self.cross_attention_block_skip_range},\n"
97
+ f" spatial_attention_timestep_skip_range={self.spatial_attention_timestep_skip_range},\n"
98
+ f" temporal_attention_timestep_skip_range={self.temporal_attention_timestep_skip_range},\n"
99
+ f" cross_attention_timestep_skip_range={self.cross_attention_timestep_skip_range},\n"
100
+ f" spatial_attention_block_identifiers={self.spatial_attention_block_identifiers},\n"
101
+ f" temporal_attention_block_identifiers={self.temporal_attention_block_identifiers},\n"
102
+ f" cross_attention_block_identifiers={self.cross_attention_block_identifiers},\n"
103
+ f" current_timestep_callback={self.current_timestep_callback}\n"
104
+ ")"
105
+ )
106
+
107
+
108
+ class PyramidAttentionBroadcastState:
109
+ r"""
110
+ State for Pyramid Attention Broadcast.
111
+
112
+ Attributes:
113
+ iteration (`int`):
114
+ The current iteration of the Pyramid Attention Broadcast. It is necessary to ensure that `reset_state` is
115
+ called before starting a new inference forward pass for PAB to work correctly.
116
+ cache (`Any`):
117
+ The cached output from the previous forward pass. This is used to re-use the attention states when the
118
+ attention computation is skipped. It is either a tensor or a tuple of tensors, depending on the module.
119
+ """
120
+
121
+ def __init__(self) -> None:
122
+ self.iteration = 0
123
+ self.cache = None
124
+
125
+ def reset(self):
126
+ self.iteration = 0
127
+ self.cache = None
128
+
129
+ def __repr__(self):
130
+ cache_repr = ""
131
+ if self.cache is None:
132
+ cache_repr = "None"
133
+ else:
134
+ cache_repr = f"Tensor(shape={self.cache.shape}, dtype={self.cache.dtype})"
135
+ return f"PyramidAttentionBroadcastState(iteration={self.iteration}, cache={cache_repr})"
136
+
137
+
138
+ class PyramidAttentionBroadcastHook(ModelHook):
139
+ r"""A hook that applies Pyramid Attention Broadcast to a given module."""
140
+
141
+ _is_stateful = True
142
+
143
+ def __init__(
144
+ self, timestep_skip_range: tuple[int, int], block_skip_range: int, current_timestep_callback: Callable[[], int]
145
+ ) -> None:
146
+ super().__init__()
147
+
148
+ self.timestep_skip_range = timestep_skip_range
149
+ self.block_skip_range = block_skip_range
150
+ self.current_timestep_callback = current_timestep_callback
151
+
152
+ def initialize_hook(self, module):
153
+ self.state = PyramidAttentionBroadcastState()
154
+ return module
155
+
156
+ def new_forward(self, module: torch.nn.Module, *args, **kwargs) -> Any:
157
+ is_within_timestep_range = (
158
+ self.timestep_skip_range[0] < self.current_timestep_callback() < self.timestep_skip_range[1]
159
+ )
160
+ should_compute_attention = (
161
+ self.state.cache is None
162
+ or self.state.iteration == 0
163
+ or not is_within_timestep_range
164
+ or self.state.iteration % self.block_skip_range == 0
165
+ )
166
+
167
+ if should_compute_attention:
168
+ output = self.fn_ref.original_forward(*args, **kwargs)
169
+ else:
170
+ output = self.state.cache
171
+
172
+ self.state.cache = output
173
+ self.state.iteration += 1
174
+ return output
175
+
176
+ def reset_state(self, module: torch.nn.Module) -> None:
177
+ self.state.reset()
178
+ return module
179
+
180
+
181
+ def apply_pyramid_attention_broadcast(module: torch.nn.Module, config: PyramidAttentionBroadcastConfig):
182
+ r"""
183
+ Apply [Pyramid Attention Broadcast](https://huggingface.co/papers/2408.12588) to a given pipeline.
184
+
185
+ PAB is an attention approximation method that leverages the similarity in attention states between timesteps to
186
+ reduce the computational cost of attention computation. The key takeaway from the paper is that the attention
187
+ similarity in the cross-attention layers between timesteps is high, followed by less similarity in the temporal and
188
+ spatial layers. This allows for the skipping of attention computation in the cross-attention layers more frequently
189
+ than in the temporal and spatial layers. Applying PAB will, therefore, speedup the inference process.
190
+
191
+ Args:
192
+ module (`torch.nn.Module`):
193
+ The module to apply Pyramid Attention Broadcast to.
194
+ config (`PyramidAttentionBroadcastConfig | None`, `optional`, defaults to `None`):
195
+ The configuration to use for Pyramid Attention Broadcast.
196
+
197
+ Example:
198
+
199
+ ```python
200
+ >>> import torch
201
+ >>> from diffusers import CogVideoXPipeline, PyramidAttentionBroadcastConfig, apply_pyramid_attention_broadcast
202
+ >>> from diffusers.utils import export_to_video
203
+
204
+ >>> pipe = CogVideoXPipeline.from_pretrained("THUDM/CogVideoX-5b", torch_dtype=torch.bfloat16)
205
+ >>> pipe.to("cuda")
206
+
207
+ >>> config = PyramidAttentionBroadcastConfig(
208
+ ... spatial_attention_block_skip_range=2,
209
+ ... spatial_attention_timestep_skip_range=(100, 800),
210
+ ... current_timestep_callback=lambda: pipe.current_timestep,
211
+ ... )
212
+ >>> apply_pyramid_attention_broadcast(pipe.transformer, config)
213
+ ```
214
+ """
215
+ if config.current_timestep_callback is None:
216
+ raise ValueError(
217
+ "The `current_timestep_callback` function must be provided in the configuration to apply Pyramid Attention Broadcast."
218
+ )
219
+
220
+ if (
221
+ config.spatial_attention_block_skip_range is None
222
+ and config.temporal_attention_block_skip_range is None
223
+ and config.cross_attention_block_skip_range is None
224
+ ):
225
+ logger.warning(
226
+ "Pyramid Attention Broadcast requires one or more of `spatial_attention_block_skip_range`, `temporal_attention_block_skip_range` "
227
+ "or `cross_attention_block_skip_range` parameters to be set to an integer, not `None`. Defaulting to using `spatial_attention_block_skip_range=2`. "
228
+ "To avoid this warning, please set one of the above parameters."
229
+ )
230
+ config.spatial_attention_block_skip_range = 2
231
+
232
+ for name, submodule in module.named_modules():
233
+ if not isinstance(submodule, (*_ATTENTION_CLASSES, AttentionModuleMixin)):
234
+ # PAB has been implemented specific to Diffusers' Attention classes. However, this does not mean that PAB
235
+ # cannot be applied to this layer. For custom layers, users can extend this functionality and implement
236
+ # their own PAB logic similar to `_apply_pyramid_attention_broadcast_on_attention_class`.
237
+ continue
238
+ _apply_pyramid_attention_broadcast_on_attention_class(name, submodule, config)
239
+
240
+
241
+ def _apply_pyramid_attention_broadcast_on_attention_class(
242
+ name: str, module: Attention, config: PyramidAttentionBroadcastConfig
243
+ ) -> bool:
244
+ is_spatial_self_attention = (
245
+ any(re.search(identifier, name) is not None for identifier in config.spatial_attention_block_identifiers)
246
+ and config.spatial_attention_block_skip_range is not None
247
+ and not getattr(module, "is_cross_attention", False)
248
+ )
249
+ is_temporal_self_attention = (
250
+ any(re.search(identifier, name) is not None for identifier in config.temporal_attention_block_identifiers)
251
+ and config.temporal_attention_block_skip_range is not None
252
+ and not getattr(module, "is_cross_attention", False)
253
+ )
254
+ is_cross_attention = (
255
+ any(re.search(identifier, name) is not None for identifier in config.cross_attention_block_identifiers)
256
+ and config.cross_attention_block_skip_range is not None
257
+ and getattr(module, "is_cross_attention", False)
258
+ )
259
+
260
+ block_skip_range, timestep_skip_range, block_type = None, None, None
261
+ if is_spatial_self_attention:
262
+ block_skip_range = config.spatial_attention_block_skip_range
263
+ timestep_skip_range = config.spatial_attention_timestep_skip_range
264
+ block_type = "spatial"
265
+ elif is_temporal_self_attention:
266
+ block_skip_range = config.temporal_attention_block_skip_range
267
+ timestep_skip_range = config.temporal_attention_timestep_skip_range
268
+ block_type = "temporal"
269
+ elif is_cross_attention:
270
+ block_skip_range = config.cross_attention_block_skip_range
271
+ timestep_skip_range = config.cross_attention_timestep_skip_range
272
+ block_type = "cross"
273
+
274
+ if block_skip_range is None or timestep_skip_range is None:
275
+ logger.info(
276
+ f'Unable to apply Pyramid Attention Broadcast to the selected layer: "{name}" because it does '
277
+ f"not match any of the required criteria for spatial, temporal or cross attention layers. Note, "
278
+ f"however, that this layer may still be valid for applying PAB. Please specify the correct "
279
+ f"block identifiers in the configuration."
280
+ )
281
+ return False
282
+
283
+ logger.debug(f"Enabling Pyramid Attention Broadcast ({block_type}) in layer: {name}")
284
+ _apply_pyramid_attention_broadcast_hook(
285
+ module, timestep_skip_range, block_skip_range, config.current_timestep_callback
286
+ )
287
+ return True
288
+
289
+
290
+ def _apply_pyramid_attention_broadcast_hook(
291
+ module: Attention | MochiAttention,
292
+ timestep_skip_range: tuple[int, int],
293
+ block_skip_range: int,
294
+ current_timestep_callback: Callable[[], int],
295
+ ):
296
+ r"""
297
+ Apply [Pyramid Attention Broadcast](https://huggingface.co/papers/2408.12588) to a given torch.nn.Module.
298
+
299
+ Args:
300
+ module (`torch.nn.Module`):
301
+ The module to apply Pyramid Attention Broadcast to.
302
+ timestep_skip_range (`tuple[int, int]`):
303
+ The range of timesteps to skip in the attention layer. The attention computations will be conditionally
304
+ skipped if the current timestep is within the specified range.
305
+ block_skip_range (`int`):
306
+ The number of times a specific attention broadcast is skipped before computing the attention states to
307
+ re-use. If this is set to the value `N`, the attention computation will be skipped `N - 1` times (i.e., old
308
+ attention states will be reused) before computing the new attention states again.
309
+ current_timestep_callback (`Callable[[], int]`):
310
+ A callback function that returns the current inference timestep.
311
+ """
312
+ registry = HookRegistry.check_if_exists_or_initialize(module)
313
+ hook = PyramidAttentionBroadcastHook(timestep_skip_range, block_skip_range, current_timestep_callback)
314
+ registry.register_hook(hook, _PYRAMID_ATTENTION_BROADCAST_HOOK)
diffusers/hooks/smoothed_energy_guidance_utils.py ADDED
@@ -0,0 +1,166 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 The HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ import math
16
+ from dataclasses import asdict, dataclass
17
+
18
+ import torch
19
+ import torch.nn.functional as F
20
+
21
+ from ..utils import get_logger
22
+ from ._common import _ALL_TRANSFORMER_BLOCK_IDENTIFIERS, _ATTENTION_CLASSES, _get_submodule_from_fqn
23
+ from .hooks import HookRegistry, ModelHook
24
+
25
+
26
+ logger = get_logger(__name__) # pylint: disable=invalid-name
27
+
28
+ _SMOOTHED_ENERGY_GUIDANCE_HOOK = "smoothed_energy_guidance_hook"
29
+
30
+
31
+ @dataclass
32
+ class SmoothedEnergyGuidanceConfig:
33
+ r"""
34
+ Configuration for skipping internal transformer blocks when executing a transformer model.
35
+
36
+ Args:
37
+ indices (`list[int]`):
38
+ The indices of the layer to skip. This is typically the first layer in the transformer block.
39
+ fqn (`str`, defaults to `"auto"`):
40
+ The fully qualified name identifying the stack of transformer blocks. Typically, this is
41
+ `transformer_blocks`, `single_transformer_blocks`, `blocks`, `layers`, or `temporal_transformer_blocks`.
42
+ For automatic detection, set this to `"auto"`. "auto" only works on DiT models. For UNet models, you must
43
+ provide the correct fqn.
44
+ _query_proj_identifiers (`list[str]`, defaults to `None`):
45
+ The identifiers for the query projection layers. Typically, these are `to_q`, `query`, or `q_proj`. If
46
+ `None`, `to_q` is used by default.
47
+ """
48
+
49
+ indices: list[int]
50
+ fqn: str = "auto"
51
+ _query_proj_identifiers: list[str] = None
52
+
53
+ def to_dict(self):
54
+ return asdict(self)
55
+
56
+ @staticmethod
57
+ def from_dict(data: dict) -> "SmoothedEnergyGuidanceConfig":
58
+ return SmoothedEnergyGuidanceConfig(**data)
59
+
60
+
61
+ class SmoothedEnergyGuidanceHook(ModelHook):
62
+ def __init__(self, blur_sigma: float = 1.0, blur_threshold_inf: float = 9999.9) -> None:
63
+ super().__init__()
64
+ self.blur_sigma = blur_sigma
65
+ self.blur_threshold_inf = blur_threshold_inf
66
+
67
+ def post_forward(self, module: torch.nn.Module, output: torch.Tensor) -> torch.Tensor:
68
+ # Copied from https://github.com/SusungHong/SEG-SDXL/blob/cf8256d640d5373541cfea3b3b6caf93272cf986/pipeline_seg.py#L172C31-L172C102
69
+ kernel_size = math.ceil(6 * self.blur_sigma) + 1 - math.ceil(6 * self.blur_sigma) % 2
70
+ smoothed_output = _gaussian_blur_2d(output, kernel_size, self.blur_sigma, self.blur_threshold_inf)
71
+ return smoothed_output
72
+
73
+
74
+ def _apply_smoothed_energy_guidance_hook(
75
+ module: torch.nn.Module, config: SmoothedEnergyGuidanceConfig, blur_sigma: float, name: str | None = None
76
+ ) -> None:
77
+ name = name or _SMOOTHED_ENERGY_GUIDANCE_HOOK
78
+
79
+ if config.fqn == "auto":
80
+ for identifier in _ALL_TRANSFORMER_BLOCK_IDENTIFIERS:
81
+ if hasattr(module, identifier):
82
+ config.fqn = identifier
83
+ break
84
+ else:
85
+ raise ValueError(
86
+ "Could not find a suitable identifier for the transformer blocks automatically. Please provide a valid "
87
+ "`fqn` (fully qualified name) that identifies a stack of transformer blocks."
88
+ )
89
+
90
+ if config._query_proj_identifiers is None:
91
+ config._query_proj_identifiers = ["to_q"]
92
+
93
+ transformer_blocks = _get_submodule_from_fqn(module, config.fqn)
94
+ blocks_found = False
95
+ for i, block in enumerate(transformer_blocks):
96
+ if i not in config.indices:
97
+ continue
98
+
99
+ blocks_found = True
100
+
101
+ for submodule_name, submodule in block.named_modules():
102
+ if not isinstance(submodule, _ATTENTION_CLASSES) or submodule.is_cross_attention:
103
+ continue
104
+ for identifier in config._query_proj_identifiers:
105
+ query_proj = getattr(submodule, identifier, None)
106
+ if query_proj is None or not isinstance(query_proj, torch.nn.Linear):
107
+ continue
108
+ logger.debug(
109
+ f"Registering smoothed energy guidance hook on {config.fqn}.{i}.{submodule_name}.{identifier}"
110
+ )
111
+ registry = HookRegistry.check_if_exists_or_initialize(query_proj)
112
+ hook = SmoothedEnergyGuidanceHook(blur_sigma)
113
+ registry.register_hook(hook, name)
114
+
115
+ if not blocks_found:
116
+ raise ValueError(
117
+ f"Could not find any transformer blocks matching the provided indices {config.indices} and "
118
+ f"fully qualified name '{config.fqn}'. Please check the indices and fqn for correctness."
119
+ )
120
+
121
+
122
+ # Modified from https://github.com/SusungHong/SEG-SDXL/blob/cf8256d640d5373541cfea3b3b6caf93272cf986/pipeline_seg.py#L71
123
+ def _gaussian_blur_2d(query: torch.Tensor, kernel_size: int, sigma: float, sigma_threshold_inf: float) -> torch.Tensor:
124
+ """
125
+ This implementation assumes that the input query is for visual (image/videos) tokens to apply the 2D gaussian blur.
126
+ However, some models use joint text-visual token attention for which this may not be suitable. Additionally, this
127
+ implementation also assumes that the visual tokens come from a square image/video. In practice, despite these
128
+ assumptions, applying the 2D square gaussian blur on the query projections generates reasonable results for
129
+ Smoothed Energy Guidance.
130
+
131
+ SEG is only supported as an experimental prototype feature for now, so the implementation may be modified in the
132
+ future without warning or guarantee of reproducibility.
133
+ """
134
+ assert query.ndim == 3
135
+
136
+ is_inf = sigma > sigma_threshold_inf
137
+ batch_size, seq_len, embed_dim = query.shape
138
+
139
+ seq_len_sqrt = int(math.sqrt(seq_len))
140
+ num_square_tokens = seq_len_sqrt * seq_len_sqrt
141
+ query_slice = query[:, :num_square_tokens, :]
142
+ query_slice = query_slice.permute(0, 2, 1)
143
+ query_slice = query_slice.reshape(batch_size, embed_dim, seq_len_sqrt, seq_len_sqrt)
144
+
145
+ if is_inf:
146
+ kernel_size = min(kernel_size, seq_len_sqrt - (seq_len_sqrt % 2 - 1))
147
+ kernel_size_half = (kernel_size - 1) / 2
148
+
149
+ x = torch.linspace(-kernel_size_half, kernel_size_half, steps=kernel_size)
150
+ pdf = torch.exp(-0.5 * (x / sigma).pow(2))
151
+ kernel1d = pdf / pdf.sum()
152
+ kernel1d = kernel1d.to(query)
153
+ kernel2d = torch.matmul(kernel1d[:, None], kernel1d[None, :])
154
+ kernel2d = kernel2d.expand(embed_dim, 1, kernel2d.shape[0], kernel2d.shape[1])
155
+
156
+ padding = [kernel_size // 2, kernel_size // 2, kernel_size // 2, kernel_size // 2]
157
+ query_slice = F.pad(query_slice, padding, mode="reflect")
158
+ query_slice = F.conv2d(query_slice, kernel2d, groups=embed_dim)
159
+ else:
160
+ query_slice[:] = query_slice.mean(dim=(-2, -1), keepdim=True)
161
+
162
+ query_slice = query_slice.reshape(batch_size, embed_dim, num_square_tokens)
163
+ query_slice = query_slice.permute(0, 2, 1)
164
+ query[:, :num_square_tokens, :] = query_slice.clone()
165
+
166
+ return query
diffusers/hooks/taylorseer_cache.py ADDED
@@ -0,0 +1,345 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import math
2
+ import re
3
+ from dataclasses import dataclass
4
+
5
+ import torch
6
+ import torch.nn as nn
7
+
8
+ from ..utils import logging
9
+ from .hooks import HookRegistry, ModelHook, StateManager
10
+
11
+
12
+ logger = logging.get_logger(__name__)
13
+ _TAYLORSEER_CACHE_HOOK = "taylorseer_cache"
14
+ _SPATIAL_ATTENTION_BLOCK_IDENTIFIERS = (
15
+ "^blocks.*attn",
16
+ "^transformer_blocks.*attn",
17
+ "^single_transformer_blocks.*attn",
18
+ )
19
+ _TEMPORAL_ATTENTION_BLOCK_IDENTIFIERS = ("^temporal_transformer_blocks.*attn",)
20
+ _TRANSFORMER_BLOCK_IDENTIFIERS = _SPATIAL_ATTENTION_BLOCK_IDENTIFIERS + _TEMPORAL_ATTENTION_BLOCK_IDENTIFIERS
21
+ _BLOCK_IDENTIFIERS = ("^[^.]*block[^.]*\\.[^.]+$",)
22
+ _PROJ_OUT_IDENTIFIERS = ("^proj_out$",)
23
+
24
+
25
+ @dataclass
26
+ class TaylorSeerCacheConfig:
27
+ """
28
+ Configuration for TaylorSeer cache. See: https://huggingface.co/papers/2503.06923
29
+
30
+ Attributes:
31
+ cache_interval (`int`, defaults to `5`):
32
+ The interval between full computation steps. After a full computation, the cached (predicted) outputs are
33
+ reused for this many subsequent denoising steps before refreshing with a new full forward pass.
34
+
35
+ disable_cache_before_step (`int`, defaults to `3`):
36
+ The denoising step index before which caching is disabled, meaning full computation is performed for the
37
+ initial steps (0 to disable_cache_before_step - 1) to gather data for Taylor series approximations. During
38
+ these steps, Taylor factors are updated, but caching/predictions are not applied. Caching begins at this
39
+ step.
40
+
41
+ disable_cache_after_step (`int`, *optional*, defaults to `None`):
42
+ The denoising step index after which caching is disabled. If set, for steps >= this value, all modules run
43
+ full computations without predictions or state updates, ensuring accuracy in later stages if needed.
44
+
45
+ max_order (`int`, defaults to `1`):
46
+ The highest order in the Taylor series expansion for approximating module outputs. Higher orders provide
47
+ better approximations but increase computation and memory usage.
48
+
49
+ taylor_factors_dtype (`torch.dtype`, defaults to `torch.bfloat16`):
50
+ Data type used for storing and computing Taylor series factors. Lower precision reduces memory but may
51
+ affect stability; higher precision improves accuracy at the cost of more memory.
52
+
53
+ skip_predict_identifiers (`list[str]`, *optional*, defaults to `None`):
54
+ Regex patterns (using `re.fullmatch`) for module names to place as "skip" in "cache" mode. In this mode,
55
+ the module computes fully during initial or refresh steps but returns a zero tensor (matching recorded
56
+ shape) during prediction steps to skip computation cheaply.
57
+
58
+ cache_identifiers (`list[str]`, *optional*, defaults to `None`):
59
+ Regex patterns (using `re.fullmatch`) for module names to place in Taylor-series caching mode, where
60
+ outputs are approximated and cached for reuse.
61
+
62
+ use_lite_mode (`bool`, *optional*, defaults to `False`):
63
+ Enables a lightweight TaylorSeer variant that minimizes memory usage by applying predefined patterns for
64
+ skipping and caching (e.g., skipping blocks and caching projections). This overrides any custom
65
+ `inactive_identifiers` or `active_identifiers`.
66
+
67
+ Notes:
68
+ - Patterns are matched using `re.fullmatch` on the module name.
69
+ - If `skip_predict_identifiers` or `cache_identifiers` are provided, only matching modules are hooked.
70
+ - If neither is provided, all attention-like modules are hooked by default.
71
+
72
+ Example of inactive and active usage:
73
+
74
+ ```py
75
+ def forward(x):
76
+ x = self.module1(x) # inactive module: returns zeros tensor based on shape recorded during full compute
77
+ x = self.module2(x) # active module: caches output here, avoiding recomputation of prior steps
78
+ return x
79
+ ```
80
+ """
81
+
82
+ cache_interval: int = 5
83
+ disable_cache_before_step: int = 3
84
+ disable_cache_after_step: int | None = None
85
+ max_order: int = 1
86
+ taylor_factors_dtype: torch.dtype | None = torch.bfloat16
87
+ skip_predict_identifiers: list[str] | None = None
88
+ cache_identifiers: list[str] | None = None
89
+ use_lite_mode: bool = False
90
+
91
+ def __repr__(self) -> str:
92
+ return (
93
+ "TaylorSeerCacheConfig("
94
+ f"cache_interval={self.cache_interval}, "
95
+ f"disable_cache_before_step={self.disable_cache_before_step}, "
96
+ f"disable_cache_after_step={self.disable_cache_after_step}, "
97
+ f"max_order={self.max_order}, "
98
+ f"taylor_factors_dtype={self.taylor_factors_dtype}, "
99
+ f"skip_predict_identifiers={self.skip_predict_identifiers}, "
100
+ f"cache_identifiers={self.cache_identifiers}, "
101
+ f"use_lite_mode={self.use_lite_mode})"
102
+ )
103
+
104
+
105
+ class TaylorSeerState:
106
+ def __init__(
107
+ self,
108
+ taylor_factors_dtype: torch.dtype | None = torch.bfloat16,
109
+ max_order: int = 1,
110
+ is_inactive: bool = False,
111
+ ):
112
+ self.taylor_factors_dtype = taylor_factors_dtype
113
+ self.max_order = max_order
114
+ self.is_inactive = is_inactive
115
+
116
+ self.module_dtypes: tuple[torch.dtype, ...] = ()
117
+ self.last_update_step: int | None = None
118
+ self.taylor_factors: dict[int, dict[int, torch.Tensor]] = {}
119
+ self.inactive_shapes: tuple[tuple[int, ...], ...] | None = None
120
+ self.device: torch.device | None = None
121
+ self.current_step: int = -1
122
+
123
+ def reset(self) -> None:
124
+ self.current_step = -1
125
+ self.last_update_step = None
126
+ self.taylor_factors = {}
127
+ self.inactive_shapes = None
128
+ self.device = None
129
+
130
+ def update(
131
+ self,
132
+ outputs: tuple[torch.Tensor, ...],
133
+ ) -> None:
134
+ self.module_dtypes = tuple(output.dtype for output in outputs)
135
+ self.device = outputs[0].device
136
+
137
+ if self.is_inactive:
138
+ self.inactive_shapes = tuple(output.shape for output in outputs)
139
+ else:
140
+ for i, features in enumerate(outputs):
141
+ new_factors: dict[int, torch.Tensor] = {0: features}
142
+ is_first_update = self.last_update_step is None
143
+ if not is_first_update:
144
+ delta_step = self.current_step - self.last_update_step
145
+ if delta_step == 0:
146
+ raise ValueError("Delta step cannot be zero for TaylorSeer update.")
147
+
148
+ # Recursive divided differences up to max_order
149
+ prev_factors = self.taylor_factors.get(i, {})
150
+ for j in range(self.max_order):
151
+ prev = prev_factors.get(j)
152
+ if prev is None:
153
+ break
154
+ new_factors[j + 1] = (new_factors[j] - prev.to(features.dtype)) / delta_step
155
+ self.taylor_factors[i] = {
156
+ order: factor.to(self.taylor_factors_dtype) for order, factor in new_factors.items()
157
+ }
158
+
159
+ self.last_update_step = self.current_step
160
+
161
+ @torch.compiler.disable
162
+ def predict(self) -> list[torch.Tensor]:
163
+ if self.last_update_step is None:
164
+ raise ValueError("Cannot predict without prior initialization/update.")
165
+
166
+ step_offset = self.current_step - self.last_update_step
167
+
168
+ outputs = []
169
+ if self.is_inactive:
170
+ if self.inactive_shapes is None:
171
+ raise ValueError("Inactive shapes not set during prediction.")
172
+ for i in range(len(self.module_dtypes)):
173
+ outputs.append(
174
+ torch.zeros(
175
+ self.inactive_shapes[i],
176
+ dtype=self.module_dtypes[i],
177
+ device=self.device,
178
+ )
179
+ )
180
+ else:
181
+ if not self.taylor_factors:
182
+ raise ValueError("Taylor factors empty during prediction.")
183
+ num_outputs = len(self.taylor_factors)
184
+ num_orders = len(self.taylor_factors[0])
185
+ for i in range(num_outputs):
186
+ output_dtype = self.module_dtypes[i]
187
+ taylor_factors = self.taylor_factors[i]
188
+ output = torch.zeros_like(taylor_factors[0], dtype=output_dtype)
189
+ for order in range(num_orders):
190
+ coeff = (step_offset**order) / math.factorial(order)
191
+ factor = taylor_factors[order]
192
+ output = output + factor.to(output_dtype) * coeff
193
+ outputs.append(output)
194
+ return outputs
195
+
196
+
197
+ class TaylorSeerCacheHook(ModelHook):
198
+ _is_stateful = True
199
+
200
+ def __init__(
201
+ self,
202
+ cache_interval: int,
203
+ disable_cache_before_step: int,
204
+ taylor_factors_dtype: torch.dtype,
205
+ state_manager: StateManager,
206
+ disable_cache_after_step: int | None = None,
207
+ ):
208
+ super().__init__()
209
+ self.cache_interval = cache_interval
210
+ self.disable_cache_before_step = disable_cache_before_step
211
+ self.disable_cache_after_step = disable_cache_after_step
212
+ self.taylor_factors_dtype = taylor_factors_dtype
213
+ self.state_manager = state_manager
214
+
215
+ def initialize_hook(self, module: torch.nn.Module):
216
+ return module
217
+
218
+ def reset_state(self, module: torch.nn.Module) -> None:
219
+ """
220
+ Reset state between sampling runs.
221
+ """
222
+ self.state_manager.reset()
223
+
224
+ @torch.compiler.disable
225
+ def _measure_should_compute(self) -> bool:
226
+ state: TaylorSeerState = self.state_manager.get_state()
227
+ state.current_step += 1
228
+ current_step = state.current_step
229
+ is_warmup_phase = current_step < self.disable_cache_before_step
230
+ is_compute_interval = (current_step - self.disable_cache_before_step - 1) % self.cache_interval == 0
231
+ is_cooldown_phase = self.disable_cache_after_step is not None and current_step >= self.disable_cache_after_step
232
+ should_compute = is_warmup_phase or is_compute_interval or is_cooldown_phase
233
+ return should_compute, state
234
+
235
+ def new_forward(self, module: torch.nn.Module, *args, **kwargs):
236
+ should_compute, state = self._measure_should_compute()
237
+ if should_compute:
238
+ outputs = self.fn_ref.original_forward(*args, **kwargs)
239
+ wrapped_outputs = (outputs,) if isinstance(outputs, torch.Tensor) else outputs
240
+ state.update(wrapped_outputs)
241
+ return outputs
242
+
243
+ outputs_list = state.predict()
244
+ return outputs_list[0] if len(outputs_list) == 1 else tuple(outputs_list)
245
+
246
+
247
+ def _resolve_patterns(config: TaylorSeerCacheConfig) -> tuple[list[str], list[str]]:
248
+ """
249
+ Resolve effective inactive and active pattern lists from config + templates.
250
+ """
251
+
252
+ inactive_patterns = config.skip_predict_identifiers if config.skip_predict_identifiers is not None else None
253
+ active_patterns = config.cache_identifiers if config.cache_identifiers is not None else None
254
+
255
+ return inactive_patterns or [], active_patterns or []
256
+
257
+
258
+ def apply_taylorseer_cache(module: torch.nn.Module, config: TaylorSeerCacheConfig):
259
+ """
260
+ Applies the TaylorSeer cache to a given pipeline (typically the transformer / UNet).
261
+
262
+ This function hooks selected modules in the model to enable caching or skipping based on the provided
263
+ configuration, reducing redundant computations in diffusion denoising loops.
264
+
265
+ Args:
266
+ module (torch.nn.Module): The model subtree to apply the hooks to.
267
+ config (TaylorSeerCacheConfig): Configuration for the cache.
268
+
269
+ Example:
270
+ ```python
271
+ >>> import torch
272
+ >>> from diffusers import FluxPipeline, TaylorSeerCacheConfig
273
+
274
+ >>> pipe = FluxPipeline.from_pretrained(
275
+ ... "black-forest-labs/FLUX.1-dev",
276
+ ... torch_dtype=torch.bfloat16,
277
+ ... )
278
+ >>> pipe.to("cuda")
279
+
280
+ >>> config = TaylorSeerCacheConfig(
281
+ ... cache_interval=5,
282
+ ... max_order=1,
283
+ ... disable_cache_before_step=3,
284
+ ... taylor_factors_dtype=torch.float32,
285
+ ... )
286
+ >>> pipe.transformer.enable_cache(config)
287
+ ```
288
+ """
289
+ inactive_patterns, active_patterns = _resolve_patterns(config)
290
+
291
+ active_patterns = active_patterns or _TRANSFORMER_BLOCK_IDENTIFIERS
292
+
293
+ if config.use_lite_mode:
294
+ logger.info("Using TaylorSeer Lite variant for cache.")
295
+ active_patterns = _PROJ_OUT_IDENTIFIERS
296
+ inactive_patterns = _BLOCK_IDENTIFIERS
297
+ if config.skip_predict_identifiers or config.cache_identifiers:
298
+ logger.warning("Lite mode overrides user patterns.")
299
+
300
+ for name, submodule in module.named_modules():
301
+ matches_inactive = any(re.fullmatch(pattern, name) for pattern in inactive_patterns)
302
+ matches_active = any(re.fullmatch(pattern, name) for pattern in active_patterns)
303
+ if not (matches_inactive or matches_active):
304
+ continue
305
+ _apply_taylorseer_cache_hook(
306
+ module=submodule,
307
+ config=config,
308
+ is_inactive=matches_inactive,
309
+ )
310
+
311
+
312
+ def _apply_taylorseer_cache_hook(
313
+ module: nn.Module,
314
+ config: TaylorSeerCacheConfig,
315
+ is_inactive: bool,
316
+ ):
317
+ """
318
+ Registers the TaylorSeer hook on the specified nn.Module.
319
+
320
+ Args:
321
+ name: Name of the module.
322
+ module: The nn.Module to be hooked.
323
+ config: Cache configuration.
324
+ is_inactive: Whether this module should operate in "inactive" mode.
325
+ """
326
+ state_manager = StateManager(
327
+ TaylorSeerState,
328
+ init_kwargs={
329
+ "taylor_factors_dtype": config.taylor_factors_dtype,
330
+ "max_order": config.max_order,
331
+ "is_inactive": is_inactive,
332
+ },
333
+ )
334
+
335
+ registry = HookRegistry.check_if_exists_or_initialize(module)
336
+
337
+ hook = TaylorSeerCacheHook(
338
+ cache_interval=config.cache_interval,
339
+ disable_cache_before_step=config.disable_cache_before_step,
340
+ taylor_factors_dtype=config.taylor_factors_dtype,
341
+ disable_cache_after_step=config.disable_cache_after_step,
342
+ state_manager=state_manager,
343
+ )
344
+
345
+ registry.register_hook(hook, _TAYLORSEER_CACHE_HOOK)
diffusers/hooks/text_kv_cache.py ADDED
@@ -0,0 +1,173 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 The HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ from dataclasses import dataclass
16
+
17
+ import torch
18
+
19
+ from .hooks import BaseState, HookRegistry, ModelHook, StateManager
20
+
21
+
22
+ _TEXT_KV_CACHE_TRANSFORMER_HOOK = "text_kv_cache_transformer"
23
+ _TEXT_KV_CACHE_BLOCK_HOOK = "text_kv_cache_block"
24
+
25
+
26
+ @dataclass
27
+ class TextKVCacheConfig:
28
+ """Enable exact (lossless) text K/V caching for transformer models.
29
+
30
+ Pre-computes per-block text key and value projections once before the denoising loop and reuses them across all
31
+ steps. Positive and negative prompts are distinguished via a stable cache key captured by a transformer-level hook
32
+ before any intermediate tensor allocations.
33
+ """
34
+
35
+ pass
36
+
37
+
38
+ class TextKVCacheState(BaseState):
39
+ """Shared state between the transformer-level and block-level hooks.
40
+
41
+ The transformer hook writes the stable ``encoder_hidden_states`` ``data_ptr()`` (captured *before* ``txt_norm``) so
42
+ that block hooks can use it as a reliable cache key across denoising steps.
43
+ """
44
+
45
+ def __init__(self):
46
+ self.key: int | None = None
47
+
48
+ def reset(self):
49
+ self.key = None
50
+
51
+
52
+ class TextKVCacheBlockState(BaseState):
53
+ """Per-block state holding cached text key/value projections."""
54
+
55
+ def __init__(self):
56
+ self.kv_cache: dict[int, tuple[torch.Tensor, torch.Tensor]] = {}
57
+
58
+ def reset(self):
59
+ self.kv_cache.clear()
60
+
61
+
62
+ class TextKVCacheTransformerHook(ModelHook):
63
+ """Captures ``encoder_hidden_states.data_ptr()`` before ``txt_norm``
64
+ and writes it to shared state for the block hooks to read."""
65
+
66
+ _is_stateful = True
67
+
68
+ def __init__(self, state_manager: StateManager):
69
+ super().__init__()
70
+ self.state_manager = state_manager
71
+
72
+ def new_forward(self, module: torch.nn.Module, *args, **kwargs):
73
+ if self.state_manager._current_context is None:
74
+ self.state_manager.set_context("inference")
75
+
76
+ encoder_hidden_states = kwargs.get("encoder_hidden_states")
77
+ if encoder_hidden_states is not None:
78
+ state: TextKVCacheState = self.state_manager.get_state()
79
+ state.key = encoder_hidden_states.data_ptr()
80
+ return self.fn_ref.original_forward(*args, **kwargs)
81
+
82
+ def reset_state(self, module: torch.nn.Module):
83
+ self.state_manager.reset()
84
+ return module
85
+
86
+
87
+ class TextKVCacheBlockHook(ModelHook):
88
+ """Caches ``(txt_key, txt_value)`` per block per unique prompt using
89
+ the stable cache key from the shared state."""
90
+
91
+ _is_stateful = True
92
+
93
+ def __init__(self, state_manager: StateManager, block_state_manager: StateManager):
94
+ super().__init__()
95
+ self.state_manager = state_manager
96
+ self.block_state_manager = block_state_manager
97
+
98
+ def new_forward(self, module: torch.nn.Module, *args, **kwargs):
99
+ from ..models.transformers.transformer_nucleusmoe_image import _apply_rotary_emb_nucleus
100
+
101
+ if self.state_manager._current_context is None:
102
+ self.state_manager.set_context("inference")
103
+
104
+ if self.block_state_manager._current_context is None:
105
+ self.block_state_manager.set_context("inference")
106
+
107
+ if "encoder_hidden_states" in kwargs:
108
+ encoder_hidden_states = kwargs["encoder_hidden_states"]
109
+ else:
110
+ encoder_hidden_states = args[1]
111
+
112
+ if "image_rotary_emb" in kwargs:
113
+ image_rotary_emb = kwargs["image_rotary_emb"]
114
+ elif len(args) > 3:
115
+ image_rotary_emb = args[3]
116
+ else:
117
+ image_rotary_emb = None
118
+
119
+ state: TextKVCacheState = self.state_manager.get_state()
120
+ cache_key = state.key
121
+
122
+ block_state: TextKVCacheBlockState = self.block_state_manager.get_state()
123
+
124
+ if cache_key not in block_state.kv_cache:
125
+ context = module.encoder_proj(encoder_hidden_states)
126
+
127
+ attn = module.attn
128
+ head_dim = attn.inner_dim // attn.heads
129
+ num_kv_heads = attn.inner_kv_dim // head_dim
130
+
131
+ txt_key = attn.add_k_proj(context).unflatten(-1, (num_kv_heads, -1))
132
+ txt_value = attn.add_v_proj(context).unflatten(-1, (num_kv_heads, -1))
133
+
134
+ if attn.norm_added_k is not None:
135
+ txt_key = attn.norm_added_k(txt_key)
136
+
137
+ if image_rotary_emb is not None:
138
+ _, txt_freqs = image_rotary_emb
139
+ txt_key = _apply_rotary_emb_nucleus(txt_key, txt_freqs, use_real=False)
140
+
141
+ block_state.kv_cache[cache_key] = (txt_key, txt_value)
142
+
143
+ txt_key, txt_value = block_state.kv_cache[cache_key]
144
+
145
+ attn_kwargs = kwargs.get("attention_kwargs") or {}
146
+ attn_kwargs["cached_txt_key"] = txt_key
147
+ attn_kwargs["cached_txt_value"] = txt_value
148
+ kwargs["attention_kwargs"] = attn_kwargs
149
+
150
+ return self.fn_ref.original_forward(*args, **kwargs)
151
+
152
+ def reset_state(self, module: torch.nn.Module):
153
+ self.block_state_manager.reset()
154
+ return module
155
+
156
+
157
+ def apply_text_kv_cache(module: torch.nn.Module, config: TextKVCacheConfig) -> None:
158
+ from ..models.transformers.transformer_nucleusmoe_image import NucleusMoEImageTransformerBlock
159
+
160
+ HookRegistry.check_if_exists_or_initialize(module)
161
+
162
+ state_manager = StateManager(TextKVCacheState)
163
+
164
+ transformer_hook = TextKVCacheTransformerHook(state_manager)
165
+ registry = HookRegistry.check_if_exists_or_initialize(module)
166
+ registry.register_hook(transformer_hook, _TEXT_KV_CACHE_TRANSFORMER_HOOK)
167
+
168
+ for _, submodule in module.named_modules():
169
+ if isinstance(submodule, NucleusMoEImageTransformerBlock):
170
+ block_state_manager = StateManager(TextKVCacheBlockState)
171
+ hook = TextKVCacheBlockHook(state_manager, block_state_manager)
172
+ block_registry = HookRegistry.check_if_exists_or_initialize(submodule)
173
+ block_registry.register_hook(hook, _TEXT_KV_CACHE_BLOCK_HOOK)
diffusers/hooks/utils.py ADDED
@@ -0,0 +1,43 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 The HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ import torch
16
+
17
+ from ._common import _ALL_TRANSFORMER_BLOCK_IDENTIFIERS, _ATTENTION_CLASSES, _FEEDFORWARD_CLASSES
18
+
19
+
20
+ def _get_identifiable_transformer_blocks_in_module(module: torch.nn.Module):
21
+ module_list_with_transformer_blocks = []
22
+ for name, submodule in module.named_modules():
23
+ name_endswith_identifier = any(name.endswith(identifier) for identifier in _ALL_TRANSFORMER_BLOCK_IDENTIFIERS)
24
+ is_ModuleList = isinstance(submodule, torch.nn.ModuleList)
25
+ if name_endswith_identifier and is_ModuleList:
26
+ module_list_with_transformer_blocks.append((name, submodule))
27
+ return module_list_with_transformer_blocks
28
+
29
+
30
+ def _get_identifiable_attention_layers_in_module(module: torch.nn.Module):
31
+ attention_layers = []
32
+ for name, submodule in module.named_modules():
33
+ if isinstance(submodule, _ATTENTION_CLASSES):
34
+ attention_layers.append((name, submodule))
35
+ return attention_layers
36
+
37
+
38
+ def _get_identifiable_feedforward_layers_in_module(module: torch.nn.Module):
39
+ feedforward_layers = []
40
+ for name, submodule in module.named_modules():
41
+ if isinstance(submodule, _FEEDFORWARD_CLASSES):
42
+ feedforward_layers.append((name, submodule))
43
+ return feedforward_layers
diffusers/image_processor.py ADDED
@@ -0,0 +1,1468 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 The HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ import math
16
+ import warnings
17
+
18
+ import numpy as np
19
+ import PIL.Image
20
+ import torch
21
+ import torch.nn.functional as F
22
+ from PIL import Image, ImageFilter, ImageOps
23
+
24
+ from .configuration_utils import ConfigMixin, register_to_config
25
+ from .utils import CONFIG_NAME, PIL_INTERPOLATION, deprecate
26
+
27
+
28
+ PipelineImageInput = (
29
+ PIL.Image.Image | np.ndarray | torch.Tensor | list[PIL.Image.Image] | list[np.ndarray] | list[torch.Tensor]
30
+ )
31
+
32
+ PipelineDepthInput = PipelineImageInput
33
+
34
+
35
+ def is_valid_image(image) -> bool:
36
+ r"""
37
+ Checks if the input is a valid image.
38
+
39
+ A valid image can be:
40
+ - A `PIL.Image.Image`.
41
+ - A 2D or 3D `np.ndarray` or `torch.Tensor` (grayscale or color image).
42
+
43
+ Args:
44
+ image (`PIL.Image.Image | np.ndarray | torch.Tensor`):
45
+ The image to validate. It can be a PIL image, a NumPy array, or a torch tensor.
46
+
47
+ Returns:
48
+ `bool`:
49
+ `True` if the input is a valid image, `False` otherwise.
50
+ """
51
+ return isinstance(image, PIL.Image.Image) or isinstance(image, (np.ndarray, torch.Tensor)) and image.ndim in (2, 3)
52
+
53
+
54
+ def is_valid_image_imagelist(images):
55
+ r"""
56
+ Checks if the input is a valid image or list of images.
57
+
58
+ The input can be one of the following formats:
59
+ - A 4D tensor or numpy array (batch of images).
60
+ - A valid single image: `PIL.Image.Image`, 2D `np.ndarray` or `torch.Tensor` (grayscale image), 3D `np.ndarray` or
61
+ `torch.Tensor`.
62
+ - A list of valid images.
63
+
64
+ Args:
65
+ images (`np.ndarray | torch.Tensor | PIL.Image.Image | list`):
66
+ The image(s) to check. Can be a batch of images (4D tensor/array), a single image, or a list of valid
67
+ images.
68
+
69
+ Returns:
70
+ `bool`:
71
+ `True` if the input is valid, `False` otherwise.
72
+ """
73
+ if isinstance(images, (np.ndarray, torch.Tensor)) and images.ndim == 4:
74
+ return True
75
+ elif is_valid_image(images):
76
+ return True
77
+ elif isinstance(images, list):
78
+ return all(is_valid_image(image) for image in images)
79
+ return False
80
+
81
+
82
+ class VaeImageProcessor(ConfigMixin):
83
+ """
84
+ Image processor for VAE.
85
+
86
+ Args:
87
+ do_resize (`bool`, *optional*, defaults to `True`):
88
+ Whether to downscale the image's (height, width) dimensions to multiples of `vae_scale_factor`. Can accept
89
+ `height` and `width` arguments from [`image_processor.VaeImageProcessor.preprocess`] method.
90
+ vae_scale_factor (`int`, *optional*, defaults to `8`):
91
+ VAE scale factor. If `do_resize` is `True`, the image is automatically resized to multiples of this factor.
92
+ resample (`str`, *optional*, defaults to `lanczos`):
93
+ Resampling filter to use when resizing the image.
94
+ do_normalize (`bool`, *optional*, defaults to `True`):
95
+ Whether to normalize the image to [-1,1].
96
+ do_binarize (`bool`, *optional*, defaults to `False`):
97
+ Whether to binarize the image to 0/1.
98
+ do_convert_rgb (`bool`, *optional*, defaults to be `False`):
99
+ Whether to convert the images to RGB format.
100
+ do_convert_grayscale (`bool`, *optional*, defaults to be `False`):
101
+ Whether to convert the images to grayscale format.
102
+ """
103
+
104
+ config_name = CONFIG_NAME
105
+
106
+ @register_to_config
107
+ def __init__(
108
+ self,
109
+ do_resize: bool = True,
110
+ vae_scale_factor: int = 8,
111
+ vae_latent_channels: int = 4,
112
+ resample: str = "lanczos",
113
+ reducing_gap: int | None = None,
114
+ do_normalize: bool = True,
115
+ do_binarize: bool = False,
116
+ do_convert_rgb: bool = False,
117
+ do_convert_grayscale: bool = False,
118
+ ):
119
+ super().__init__()
120
+ if do_convert_rgb and do_convert_grayscale:
121
+ raise ValueError(
122
+ "`do_convert_rgb` and `do_convert_grayscale` can not both be set to `True`,"
123
+ " if you intended to convert the image into RGB format, please set `do_convert_grayscale = False`.",
124
+ " if you intended to convert the image into grayscale format, please set `do_convert_rgb = False`",
125
+ )
126
+
127
+ @staticmethod
128
+ def numpy_to_pil(images: np.ndarray) -> list[PIL.Image.Image]:
129
+ r"""
130
+ Convert a numpy image or a batch of images to a PIL image.
131
+
132
+ Args:
133
+ images (`np.ndarray`):
134
+ The image array to convert to PIL format.
135
+
136
+ Returns:
137
+ `list[PIL.Image.Image]`:
138
+ A list of PIL images.
139
+ """
140
+ if images.ndim == 3:
141
+ images = images[None, ...]
142
+ images = (images * 255).round().astype("uint8")
143
+ if images.shape[-1] == 1:
144
+ # special case for grayscale (single channel) images
145
+ pil_images = [Image.fromarray(image.squeeze(), mode="L") for image in images]
146
+ else:
147
+ pil_images = [Image.fromarray(image) for image in images]
148
+
149
+ return pil_images
150
+
151
+ @staticmethod
152
+ def pil_to_numpy(images: list[PIL.Image.Image] | PIL.Image.Image) -> np.ndarray:
153
+ r"""
154
+ Convert a PIL image or a list of PIL images to NumPy arrays.
155
+
156
+ Args:
157
+ images (`PIL.Image.Image` or `list[PIL.Image.Image]`):
158
+ The PIL image or list of images to convert to NumPy format.
159
+
160
+ Returns:
161
+ `np.ndarray`:
162
+ A NumPy array representation of the images.
163
+ """
164
+ if not isinstance(images, list):
165
+ images = [images]
166
+ images = [np.array(image).astype(np.float32) / 255.0 for image in images]
167
+ images = np.stack(images, axis=0)
168
+
169
+ return images
170
+
171
+ @staticmethod
172
+ def numpy_to_pt(images: np.ndarray) -> torch.Tensor:
173
+ r"""
174
+ Convert a NumPy image to a PyTorch tensor.
175
+
176
+ Args:
177
+ images (`np.ndarray`):
178
+ The NumPy image array to convert to PyTorch format.
179
+
180
+ Returns:
181
+ `torch.Tensor`:
182
+ A PyTorch tensor representation of the images.
183
+ """
184
+ if images.ndim == 3:
185
+ images = images[..., None]
186
+
187
+ images = torch.from_numpy(images.transpose(0, 3, 1, 2))
188
+ return images
189
+
190
+ @staticmethod
191
+ def pt_to_numpy(images: torch.Tensor) -> np.ndarray:
192
+ r"""
193
+ Convert a PyTorch tensor to a NumPy image.
194
+
195
+ Args:
196
+ images (`torch.Tensor`):
197
+ The PyTorch tensor to convert to NumPy format.
198
+
199
+ Returns:
200
+ `np.ndarray`:
201
+ A NumPy array representation of the images.
202
+ """
203
+ images = images.cpu().permute(0, 2, 3, 1).float().numpy()
204
+ return images
205
+
206
+ @staticmethod
207
+ def normalize(images: np.ndarray | torch.Tensor) -> np.ndarray | torch.Tensor:
208
+ r"""
209
+ Normalize an image array to [-1,1].
210
+
211
+ Args:
212
+ images (`np.ndarray` or `torch.Tensor`):
213
+ The image array to normalize.
214
+
215
+ Returns:
216
+ `np.ndarray` or `torch.Tensor`:
217
+ The normalized image array.
218
+ """
219
+ return 2.0 * images - 1.0
220
+
221
+ @staticmethod
222
+ def denormalize(images: np.ndarray | torch.Tensor) -> np.ndarray | torch.Tensor:
223
+ r"""
224
+ Denormalize an image array to [0,1].
225
+
226
+ Args:
227
+ images (`np.ndarray` or `torch.Tensor`):
228
+ The image array to denormalize.
229
+
230
+ Returns:
231
+ `np.ndarray` or `torch.Tensor`:
232
+ The denormalized image array.
233
+ """
234
+ return (images * 0.5 + 0.5).clamp(0, 1)
235
+
236
+ @staticmethod
237
+ def convert_to_rgb(image: PIL.Image.Image) -> PIL.Image.Image:
238
+ r"""
239
+ Converts a PIL image to RGB format.
240
+
241
+ Args:
242
+ image (`PIL.Image.Image`):
243
+ The PIL image to convert to RGB.
244
+
245
+ Returns:
246
+ `PIL.Image.Image`:
247
+ The RGB-converted PIL image.
248
+ """
249
+ image = image.convert("RGB")
250
+
251
+ return image
252
+
253
+ @staticmethod
254
+ def convert_to_grayscale(image: PIL.Image.Image) -> PIL.Image.Image:
255
+ r"""
256
+ Converts a given PIL image to grayscale.
257
+
258
+ Args:
259
+ image (`PIL.Image.Image`):
260
+ The input image to convert.
261
+
262
+ Returns:
263
+ `PIL.Image.Image`:
264
+ The image converted to grayscale.
265
+ """
266
+ image = image.convert("L")
267
+
268
+ return image
269
+
270
+ @staticmethod
271
+ def blur(image: PIL.Image.Image, blur_factor: int = 4) -> PIL.Image.Image:
272
+ r"""
273
+ Applies Gaussian blur to an image.
274
+
275
+ Args:
276
+ image (`PIL.Image.Image`):
277
+ The PIL image to convert to grayscale.
278
+
279
+ Returns:
280
+ `PIL.Image.Image`:
281
+ The grayscale-converted PIL image.
282
+ """
283
+ image = image.filter(ImageFilter.GaussianBlur(blur_factor))
284
+
285
+ return image
286
+
287
+ @staticmethod
288
+ def get_crop_region(mask_image: PIL.Image.Image, width: int, height: int, pad=0):
289
+ r"""
290
+ Finds a rectangular region that contains all masked ares in an image, and expands region to match the aspect
291
+ ratio of the original image; for example, if user drew mask in a 128x32 region, and the dimensions for
292
+ processing are 512x512, the region will be expanded to 128x128.
293
+
294
+ Args:
295
+ mask_image (PIL.Image.Image): Mask image.
296
+ width (int): Width of the image to be processed.
297
+ height (int): Height of the image to be processed.
298
+ pad (int, optional): Padding to be added to the crop region. Defaults to 0.
299
+
300
+ Returns:
301
+ tuple: (x1, y1, x2, y2) represent a rectangular region that contains all masked ares in an image and
302
+ matches the original aspect ratio.
303
+ """
304
+
305
+ mask_image = mask_image.convert("L")
306
+ mask = np.array(mask_image)
307
+
308
+ # 1. find a rectangular region that contains all masked ares in an image
309
+ h, w = mask.shape
310
+ crop_left = 0
311
+ for i in range(w):
312
+ if not (mask[:, i] == 0).all():
313
+ break
314
+ crop_left += 1
315
+
316
+ crop_right = 0
317
+ for i in reversed(range(w)):
318
+ if not (mask[:, i] == 0).all():
319
+ break
320
+ crop_right += 1
321
+
322
+ crop_top = 0
323
+ for i in range(h):
324
+ if not (mask[i] == 0).all():
325
+ break
326
+ crop_top += 1
327
+
328
+ crop_bottom = 0
329
+ for i in reversed(range(h)):
330
+ if not (mask[i] == 0).all():
331
+ break
332
+ crop_bottom += 1
333
+
334
+ # 2. add padding to the crop region
335
+ x1, y1, x2, y2 = (
336
+ int(max(crop_left - pad, 0)),
337
+ int(max(crop_top - pad, 0)),
338
+ int(min(w - crop_right + pad, w)),
339
+ int(min(h - crop_bottom + pad, h)),
340
+ )
341
+
342
+ # 3. expands crop region to match the aspect ratio of the image to be processed
343
+ ratio_crop_region = (x2 - x1) / (y2 - y1)
344
+ ratio_processing = width / height
345
+
346
+ if ratio_crop_region > ratio_processing:
347
+ desired_height = (x2 - x1) / ratio_processing
348
+ desired_height_diff = int(desired_height - (y2 - y1))
349
+ y1 -= desired_height_diff // 2
350
+ y2 += desired_height_diff - desired_height_diff // 2
351
+ if y2 >= mask_image.height:
352
+ diff = y2 - mask_image.height
353
+ y2 -= diff
354
+ y1 -= diff
355
+ if y1 < 0:
356
+ y2 -= y1
357
+ y1 -= y1
358
+ if y2 >= mask_image.height:
359
+ y2 = mask_image.height
360
+ else:
361
+ desired_width = (y2 - y1) * ratio_processing
362
+ desired_width_diff = int(desired_width - (x2 - x1))
363
+ x1 -= desired_width_diff // 2
364
+ x2 += desired_width_diff - desired_width_diff // 2
365
+ if x2 >= mask_image.width:
366
+ diff = x2 - mask_image.width
367
+ x2 -= diff
368
+ x1 -= diff
369
+ if x1 < 0:
370
+ x2 -= x1
371
+ x1 -= x1
372
+ if x2 >= mask_image.width:
373
+ x2 = mask_image.width
374
+
375
+ return x1, y1, x2, y2
376
+
377
+ def _resize_and_fill(
378
+ self,
379
+ image: PIL.Image.Image,
380
+ width: int,
381
+ height: int,
382
+ ) -> PIL.Image.Image:
383
+ r"""
384
+ Resize the image to fit within the specified width and height, maintaining the aspect ratio, and then center
385
+ the image within the dimensions, filling empty with data from image.
386
+
387
+ Args:
388
+ image (`PIL.Image.Image`):
389
+ The image to resize and fill.
390
+ width (`int`):
391
+ The width to resize the image to.
392
+ height (`int`):
393
+ The height to resize the image to.
394
+
395
+ Returns:
396
+ `PIL.Image.Image`:
397
+ The resized and filled image.
398
+ """
399
+
400
+ ratio = width / height
401
+ src_ratio = image.width / image.height
402
+
403
+ src_w = width if ratio < src_ratio else image.width * height // image.height
404
+ src_h = height if ratio >= src_ratio else image.height * width // image.width
405
+
406
+ resized = image.resize((src_w, src_h), resample=PIL_INTERPOLATION[self.config.resample])
407
+ res = Image.new("RGB", (width, height))
408
+ res.paste(resized, box=(width // 2 - src_w // 2, height // 2 - src_h // 2))
409
+
410
+ if ratio < src_ratio:
411
+ fill_height = height // 2 - src_h // 2
412
+ if fill_height > 0:
413
+ res.paste(resized.resize((width, fill_height), box=(0, 0, width, 0)), box=(0, 0))
414
+ res.paste(
415
+ resized.resize((width, fill_height), box=(0, resized.height, width, resized.height)),
416
+ box=(0, fill_height + src_h),
417
+ )
418
+ elif ratio > src_ratio:
419
+ fill_width = width // 2 - src_w // 2
420
+ if fill_width > 0:
421
+ res.paste(resized.resize((fill_width, height), box=(0, 0, 0, height)), box=(0, 0))
422
+ res.paste(
423
+ resized.resize((fill_width, height), box=(resized.width, 0, resized.width, height)),
424
+ box=(fill_width + src_w, 0),
425
+ )
426
+
427
+ return res
428
+
429
+ def _resize_and_crop(
430
+ self,
431
+ image: PIL.Image.Image,
432
+ width: int,
433
+ height: int,
434
+ ) -> PIL.Image.Image:
435
+ r"""
436
+ Resize the image to fit within the specified width and height, maintaining the aspect ratio, and then center
437
+ the image within the dimensions, cropping the excess.
438
+
439
+ Args:
440
+ image (`PIL.Image.Image`):
441
+ The image to resize and crop.
442
+ width (`int`):
443
+ The width to resize the image to.
444
+ height (`int`):
445
+ The height to resize the image to.
446
+
447
+ Returns:
448
+ `PIL.Image.Image`:
449
+ The resized and cropped image.
450
+ """
451
+ ratio = width / height
452
+ src_ratio = image.width / image.height
453
+
454
+ src_w = width if ratio > src_ratio else image.width * height // image.height
455
+ src_h = height if ratio <= src_ratio else image.height * width // image.width
456
+
457
+ resized = image.resize((src_w, src_h), resample=PIL_INTERPOLATION[self.config.resample])
458
+ res = Image.new("RGB", (width, height))
459
+ res.paste(resized, box=(width // 2 - src_w // 2, height // 2 - src_h // 2))
460
+ return res
461
+
462
+ def resize(
463
+ self,
464
+ image: PIL.Image.Image | np.ndarray | torch.Tensor,
465
+ height: int,
466
+ width: int,
467
+ resize_mode: str = "default", # "default", "fill", "crop"
468
+ ) -> PIL.Image.Image | np.ndarray | torch.Tensor:
469
+ """
470
+ Resize image.
471
+
472
+ Args:
473
+ image (`PIL.Image.Image`, `np.ndarray` or `torch.Tensor`):
474
+ The image input, can be a PIL image, numpy array or pytorch tensor.
475
+ height (`int`):
476
+ The height to resize to.
477
+ width (`int`):
478
+ The width to resize to.
479
+ resize_mode (`str`, *optional*, defaults to `default`):
480
+ The resize mode to use, can be one of `default` or `fill`. If `default`, will resize the image to fit
481
+ within the specified width and height, and it may not maintaining the original aspect ratio. If `fill`,
482
+ will resize the image to fit within the specified width and height, maintaining the aspect ratio, and
483
+ then center the image within the dimensions, filling empty with data from image. If `crop`, will resize
484
+ the image to fit within the specified width and height, maintaining the aspect ratio, and then center
485
+ the image within the dimensions, cropping the excess. Note that resize_mode `fill` and `crop` are only
486
+ supported for PIL image input.
487
+
488
+ Returns:
489
+ `PIL.Image.Image`, `np.ndarray` or `torch.Tensor`:
490
+ The resized image.
491
+ """
492
+ if resize_mode != "default" and not isinstance(image, PIL.Image.Image):
493
+ raise ValueError(f"Only PIL image input is supported for resize_mode {resize_mode}")
494
+ if isinstance(image, PIL.Image.Image):
495
+ if resize_mode == "default":
496
+ image = image.resize(
497
+ (width, height),
498
+ resample=PIL_INTERPOLATION[self.config.resample],
499
+ reducing_gap=self.config.reducing_gap,
500
+ )
501
+ elif resize_mode == "fill":
502
+ image = self._resize_and_fill(image, width, height)
503
+ elif resize_mode == "crop":
504
+ image = self._resize_and_crop(image, width, height)
505
+ else:
506
+ raise ValueError(f"resize_mode {resize_mode} is not supported")
507
+
508
+ elif isinstance(image, torch.Tensor):
509
+ image = torch.nn.functional.interpolate(
510
+ image,
511
+ size=(height, width),
512
+ )
513
+ elif isinstance(image, np.ndarray):
514
+ image = self.numpy_to_pt(image)
515
+ image = torch.nn.functional.interpolate(
516
+ image,
517
+ size=(height, width),
518
+ )
519
+ image = self.pt_to_numpy(image)
520
+
521
+ return image
522
+
523
+ def binarize(self, image: PIL.Image.Image) -> PIL.Image.Image:
524
+ """
525
+ Create a mask.
526
+
527
+ Args:
528
+ image (`PIL.Image.Image`):
529
+ The image input, should be a PIL image.
530
+
531
+ Returns:
532
+ `PIL.Image.Image`:
533
+ The binarized image. Values less than 0.5 are set to 0, values greater than 0.5 are set to 1.
534
+ """
535
+ image[image < 0.5] = 0
536
+ image[image >= 0.5] = 1
537
+
538
+ return image
539
+
540
+ def _denormalize_conditionally(
541
+ self, images: torch.Tensor, do_denormalize: list[bool] | None = None
542
+ ) -> torch.Tensor:
543
+ r"""
544
+ Denormalize a batch of images based on a condition list.
545
+
546
+ Args:
547
+ images (`torch.Tensor`):
548
+ The input image tensor.
549
+ do_denormalize (`Optional[list[bool]`, *optional*, defaults to `None`):
550
+ A list of booleans indicating whether to denormalize each image in the batch. If `None`, will use the
551
+ value of `do_normalize` in the `VaeImageProcessor` config.
552
+ """
553
+ if do_denormalize is None:
554
+ return self.denormalize(images) if self.config.do_normalize else images
555
+
556
+ return torch.stack(
557
+ [self.denormalize(images[i]) if do_denormalize[i] else images[i] for i in range(images.shape[0])]
558
+ )
559
+
560
+ def get_default_height_width(
561
+ self,
562
+ image: PIL.Image.Image | np.ndarray | torch.Tensor,
563
+ height: int | None = None,
564
+ width: int | None = None,
565
+ ) -> tuple[int, int]:
566
+ r"""
567
+ Returns the height and width of the image, downscaled to the next integer multiple of `vae_scale_factor`.
568
+
569
+ Args:
570
+ image (`PIL.Image.Image | np.ndarray | torch.Tensor`):
571
+ The image input, which can be a PIL image, NumPy array, or PyTorch tensor. If it is a NumPy array, it
572
+ should have shape `[batch, height, width]` or `[batch, height, width, channels]`. If it is a PyTorch
573
+ tensor, it should have shape `[batch, channels, height, width]`.
574
+ height (`int | None`, *optional*, defaults to `None`):
575
+ The height of the preprocessed image. If `None`, the height of the `image` input will be used.
576
+ width (`int | None`, *optional*, defaults to `None`):
577
+ The width of the preprocessed image. If `None`, the width of the `image` input will be used.
578
+
579
+ Returns:
580
+ `tuple[int, int]`:
581
+ A tuple containing the height and width, both resized to the nearest integer multiple of
582
+ `vae_scale_factor`.
583
+ """
584
+
585
+ if height is None:
586
+ if isinstance(image, PIL.Image.Image):
587
+ height = image.height
588
+ elif isinstance(image, torch.Tensor):
589
+ height = image.shape[2]
590
+ else:
591
+ height = image.shape[1]
592
+
593
+ if width is None:
594
+ if isinstance(image, PIL.Image.Image):
595
+ width = image.width
596
+ elif isinstance(image, torch.Tensor):
597
+ width = image.shape[3]
598
+ else:
599
+ width = image.shape[2]
600
+
601
+ width, height = (
602
+ x - x % self.config.vae_scale_factor for x in (width, height)
603
+ ) # resize to integer multiple of vae_scale_factor
604
+
605
+ return height, width
606
+
607
+ def preprocess(
608
+ self,
609
+ image: PipelineImageInput,
610
+ height: int | None = None,
611
+ width: int | None = None,
612
+ resize_mode: str = "default", # "default", "fill", "crop"
613
+ crops_coords: tuple[int, int, int, int] | None = None,
614
+ ) -> torch.Tensor:
615
+ """
616
+ Preprocess the image input.
617
+
618
+ Args:
619
+ image (`PipelineImageInput`):
620
+ The image input, accepted formats are PIL images, NumPy arrays, PyTorch tensors; Also accept list of
621
+ supported formats.
622
+ height (`int`, *optional*):
623
+ The height in preprocessed image. If `None`, will use the `get_default_height_width()` to get default
624
+ height.
625
+ width (`int`, *optional*):
626
+ The width in preprocessed. If `None`, will use get_default_height_width()` to get the default width.
627
+ resize_mode (`str`, *optional*, defaults to `default`):
628
+ The resize mode, can be one of `default` or `fill`. If `default`, will resize the image to fit within
629
+ the specified width and height, and it may not maintaining the original aspect ratio. If `fill`, will
630
+ resize the image to fit within the specified width and height, maintaining the aspect ratio, and then
631
+ center the image within the dimensions, filling empty with data from image. If `crop`, will resize the
632
+ image to fit within the specified width and height, maintaining the aspect ratio, and then center the
633
+ image within the dimensions, cropping the excess. Note that resize_mode `fill` and `crop` are only
634
+ supported for PIL image input.
635
+ crops_coords (`list[tuple[int, int, int, int]]`, *optional*, defaults to `None`):
636
+ The crop coordinates for each image in the batch. If `None`, will not crop the image.
637
+
638
+ Returns:
639
+ `torch.Tensor`:
640
+ The preprocessed image.
641
+ """
642
+ supported_formats = (PIL.Image.Image, np.ndarray, torch.Tensor)
643
+
644
+ # Expand the missing dimension for 3-dimensional pytorch tensor or numpy array that represents grayscale image
645
+ if self.config.do_convert_grayscale and isinstance(image, (torch.Tensor, np.ndarray)) and image.ndim == 3:
646
+ if isinstance(image, torch.Tensor):
647
+ # if image is a pytorch tensor could have 2 possible shapes:
648
+ # 1. batch x height x width: we should insert the channel dimension at position 1
649
+ # 2. channel x height x width: we should insert batch dimension at position 0,
650
+ # however, since both channel and batch dimension has same size 1, it is same to insert at position 1
651
+ # for simplicity, we insert a dimension of size 1 at position 1 for both cases
652
+ image = image.unsqueeze(1)
653
+ else:
654
+ # if it is a numpy array, it could have 2 possible shapes:
655
+ # 1. batch x height x width: insert channel dimension on last position
656
+ # 2. height x width x channel: insert batch dimension on first position
657
+ if image.shape[-1] == 1:
658
+ image = np.expand_dims(image, axis=0)
659
+ else:
660
+ image = np.expand_dims(image, axis=-1)
661
+
662
+ if isinstance(image, list) and isinstance(image[0], np.ndarray) and image[0].ndim == 4:
663
+ warnings.warn(
664
+ "Passing `image` as a list of 4d np.ndarray is deprecated."
665
+ "Please concatenate the list along the batch dimension and pass it as a single 4d np.ndarray",
666
+ FutureWarning,
667
+ )
668
+ image = np.concatenate(image, axis=0)
669
+ if isinstance(image, list) and isinstance(image[0], torch.Tensor) and image[0].ndim == 4:
670
+ warnings.warn(
671
+ "Passing `image` as a list of 4d torch.Tensor is deprecated."
672
+ "Please concatenate the list along the batch dimension and pass it as a single 4d torch.Tensor",
673
+ FutureWarning,
674
+ )
675
+ image = torch.cat(image, axis=0)
676
+
677
+ if not is_valid_image_imagelist(image):
678
+ raise ValueError(
679
+ f"Input is in incorrect format. Currently, we only support {', '.join(str(x) for x in supported_formats)}"
680
+ )
681
+ if not isinstance(image, list):
682
+ image = [image]
683
+
684
+ if isinstance(image[0], PIL.Image.Image):
685
+ if crops_coords is not None:
686
+ image = [i.crop(crops_coords) for i in image]
687
+ if self.config.do_resize:
688
+ height, width = self.get_default_height_width(image[0], height, width)
689
+ image = [self.resize(i, height, width, resize_mode=resize_mode) for i in image]
690
+ if self.config.do_convert_rgb:
691
+ image = [self.convert_to_rgb(i) for i in image]
692
+ elif self.config.do_convert_grayscale:
693
+ image = [self.convert_to_grayscale(i) for i in image]
694
+ image = self.pil_to_numpy(image) # to np
695
+ image = self.numpy_to_pt(image) # to pt
696
+
697
+ elif isinstance(image[0], np.ndarray):
698
+ image = np.concatenate(image, axis=0) if image[0].ndim == 4 else np.stack(image, axis=0)
699
+
700
+ image = self.numpy_to_pt(image)
701
+
702
+ height, width = self.get_default_height_width(image, height, width)
703
+ if self.config.do_resize:
704
+ image = self.resize(image, height, width)
705
+
706
+ elif isinstance(image[0], torch.Tensor):
707
+ image = torch.cat(image, axis=0) if image[0].ndim == 4 else torch.stack(image, axis=0)
708
+
709
+ if self.config.do_convert_grayscale and image.ndim == 3:
710
+ image = image.unsqueeze(1)
711
+
712
+ channel = image.shape[1]
713
+ # don't need any preprocess if the image is latents
714
+ if channel == self.config.vae_latent_channels:
715
+ return image
716
+
717
+ height, width = self.get_default_height_width(image, height, width)
718
+ if self.config.do_resize:
719
+ image = self.resize(image, height, width)
720
+
721
+ # expected range [0,1], normalize to [-1,1]
722
+ do_normalize = self.config.do_normalize
723
+ if do_normalize and image.min() < 0:
724
+ warnings.warn(
725
+ "Passing `image` as torch tensor with value range in [-1,1] is deprecated. The expected value range for image tensor is [0,1] "
726
+ f"when passing as pytorch tensor or numpy Array. You passed `image` with value range [{image.min()},{image.max()}]",
727
+ FutureWarning,
728
+ )
729
+ do_normalize = False
730
+ if do_normalize:
731
+ image = self.normalize(image)
732
+
733
+ if self.config.do_binarize:
734
+ image = self.binarize(image)
735
+
736
+ return image
737
+
738
+ def postprocess(
739
+ self,
740
+ image: torch.Tensor,
741
+ output_type: str = "pil",
742
+ do_denormalize: list[bool] | None = None,
743
+ ) -> PIL.Image.Image | np.ndarray | torch.Tensor:
744
+ """
745
+ Postprocess the image output from tensor to `output_type`.
746
+
747
+ Args:
748
+ image (`torch.Tensor`):
749
+ The image input, should be a pytorch tensor with shape `B x C x H x W`.
750
+ output_type (`str`, *optional*, defaults to `pil`):
751
+ The output type of the image, can be one of `pil`, `np`, `pt`, `latent`.
752
+ do_denormalize (`list[bool]`, *optional*, defaults to `None`):
753
+ Whether to denormalize the image to [0,1]. If `None`, will use the value of `do_normalize` in the
754
+ `VaeImageProcessor` config.
755
+
756
+ Returns:
757
+ `PIL.Image.Image`, `np.ndarray` or `torch.Tensor`:
758
+ The postprocessed image.
759
+ """
760
+ if not isinstance(image, torch.Tensor):
761
+ raise ValueError(
762
+ f"Input for postprocessing is in incorrect format: {type(image)}. We only support pytorch tensor"
763
+ )
764
+ if output_type not in ["latent", "pt", "np", "pil"]:
765
+ deprecation_message = (
766
+ f"the output_type {output_type} is outdated and has been set to `np`. Please make sure to set it to one of these instead: "
767
+ "`pil`, `np`, `pt`, `latent`"
768
+ )
769
+ deprecate("Unsupported output_type", "1.0.0", deprecation_message, standard_warn=False)
770
+ output_type = "np"
771
+
772
+ if output_type == "latent":
773
+ return image
774
+
775
+ image = self._denormalize_conditionally(image, do_denormalize)
776
+
777
+ if output_type == "pt":
778
+ return image
779
+
780
+ image = self.pt_to_numpy(image)
781
+
782
+ if output_type == "np":
783
+ return image
784
+
785
+ if output_type == "pil":
786
+ return self.numpy_to_pil(image)
787
+
788
+ def apply_overlay(
789
+ self,
790
+ mask: PIL.Image.Image,
791
+ init_image: PIL.Image.Image,
792
+ image: PIL.Image.Image,
793
+ crop_coords: tuple[int, int, int, int] | None = None,
794
+ ) -> PIL.Image.Image:
795
+ r"""
796
+ Applies an overlay of the mask and the inpainted image on the original image.
797
+
798
+ Args:
799
+ mask (`PIL.Image.Image`):
800
+ The mask image that highlights regions to overlay.
801
+ init_image (`PIL.Image.Image`):
802
+ The original image to which the overlay is applied.
803
+ image (`PIL.Image.Image`):
804
+ The image to overlay onto the original.
805
+ crop_coords (`tuple[int, int, int, int]`, *optional*):
806
+ Coordinates to crop the image. If provided, the image will be cropped accordingly.
807
+
808
+ Returns:
809
+ `PIL.Image.Image`:
810
+ The final image with the overlay applied.
811
+ """
812
+
813
+ width, height = init_image.width, init_image.height
814
+
815
+ init_image_masked = PIL.Image.new("RGBa", (width, height))
816
+ init_image_masked.paste(init_image.convert("RGBA").convert("RGBa"), mask=ImageOps.invert(mask.convert("L")))
817
+
818
+ init_image_masked = init_image_masked.convert("RGBA")
819
+
820
+ if crop_coords is not None:
821
+ x, y, x2, y2 = crop_coords
822
+ w = x2 - x
823
+ h = y2 - y
824
+ base_image = PIL.Image.new("RGBA", (width, height))
825
+ image = self.resize(image, height=h, width=w, resize_mode="crop")
826
+ base_image.paste(image, (x, y))
827
+ image = base_image.convert("RGB")
828
+
829
+ image = image.convert("RGBA")
830
+ image.alpha_composite(init_image_masked)
831
+ image = image.convert("RGB")
832
+
833
+ return image
834
+
835
+
836
+ class InpaintProcessor(ConfigMixin):
837
+ """
838
+ Image processor for inpainting image and mask.
839
+ """
840
+
841
+ config_name = CONFIG_NAME
842
+
843
+ @register_to_config
844
+ def __init__(
845
+ self,
846
+ do_resize: bool = True,
847
+ vae_scale_factor: int = 8,
848
+ vae_latent_channels: int = 4,
849
+ resample: str = "lanczos",
850
+ reducing_gap: int | None = None,
851
+ do_normalize: bool = True,
852
+ do_binarize: bool = False,
853
+ do_convert_grayscale: bool = False,
854
+ mask_do_normalize: bool = False,
855
+ mask_do_binarize: bool = True,
856
+ mask_do_convert_grayscale: bool = True,
857
+ ):
858
+ super().__init__()
859
+
860
+ self._image_processor = VaeImageProcessor(
861
+ do_resize=do_resize,
862
+ vae_scale_factor=vae_scale_factor,
863
+ vae_latent_channels=vae_latent_channels,
864
+ resample=resample,
865
+ reducing_gap=reducing_gap,
866
+ do_normalize=do_normalize,
867
+ do_binarize=do_binarize,
868
+ do_convert_grayscale=do_convert_grayscale,
869
+ )
870
+ self._mask_processor = VaeImageProcessor(
871
+ do_resize=do_resize,
872
+ vae_scale_factor=vae_scale_factor,
873
+ vae_latent_channels=vae_latent_channels,
874
+ resample=resample,
875
+ reducing_gap=reducing_gap,
876
+ do_normalize=mask_do_normalize,
877
+ do_binarize=mask_do_binarize,
878
+ do_convert_grayscale=mask_do_convert_grayscale,
879
+ )
880
+
881
+ def preprocess(
882
+ self,
883
+ image: PIL.Image.Image,
884
+ mask: PIL.Image.Image | None = None,
885
+ height: int | None = None,
886
+ width: int | None = None,
887
+ padding_mask_crop: int | None = None,
888
+ ) -> tuple[torch.Tensor, torch.Tensor]:
889
+ """
890
+ Preprocess the image and mask.
891
+ """
892
+ if mask is None and padding_mask_crop is not None:
893
+ raise ValueError("mask must be provided if padding_mask_crop is provided")
894
+
895
+ # if mask is None, same behavior as regular image processor
896
+ if mask is None:
897
+ return self._image_processor.preprocess(image, height=height, width=width)
898
+
899
+ if padding_mask_crop is not None:
900
+ crops_coords = self._image_processor.get_crop_region(mask, width, height, pad=padding_mask_crop)
901
+ resize_mode = "fill"
902
+ else:
903
+ crops_coords = None
904
+ resize_mode = "default"
905
+
906
+ processed_image = self._image_processor.preprocess(
907
+ image,
908
+ height=height,
909
+ width=width,
910
+ crops_coords=crops_coords,
911
+ resize_mode=resize_mode,
912
+ )
913
+
914
+ processed_mask = self._mask_processor.preprocess(
915
+ mask,
916
+ height=height,
917
+ width=width,
918
+ resize_mode=resize_mode,
919
+ crops_coords=crops_coords,
920
+ )
921
+
922
+ if crops_coords is not None:
923
+ postprocessing_kwargs = {
924
+ "crops_coords": crops_coords,
925
+ "original_image": image,
926
+ "original_mask": mask,
927
+ }
928
+ else:
929
+ postprocessing_kwargs = {
930
+ "crops_coords": None,
931
+ "original_image": None,
932
+ "original_mask": None,
933
+ }
934
+
935
+ return processed_image, processed_mask, postprocessing_kwargs
936
+
937
+ def postprocess(
938
+ self,
939
+ image: torch.Tensor,
940
+ output_type: str = "pil",
941
+ original_image: PIL.Image.Image | None = None,
942
+ original_mask: PIL.Image.Image | None = None,
943
+ crops_coords: tuple[int, int, int, int] | None = None,
944
+ ) -> tuple[PIL.Image.Image, PIL.Image.Image]:
945
+ """
946
+ Postprocess the image, optionally apply mask overlay
947
+ """
948
+ image = self._image_processor.postprocess(
949
+ image,
950
+ output_type=output_type,
951
+ )
952
+ # optionally apply the mask overlay
953
+ if crops_coords is not None and (original_image is None or original_mask is None):
954
+ raise ValueError("original_image and original_mask must be provided if crops_coords is provided")
955
+
956
+ elif crops_coords is not None and output_type != "pil":
957
+ raise ValueError("output_type must be 'pil' if crops_coords is provided")
958
+
959
+ elif crops_coords is not None:
960
+ image = [
961
+ self._image_processor.apply_overlay(original_mask, original_image, i, crops_coords) for i in image
962
+ ]
963
+
964
+ return image
965
+
966
+
967
+ class VaeImageProcessorLDM3D(VaeImageProcessor):
968
+ """
969
+ Image processor for VAE LDM3D.
970
+
971
+ Args:
972
+ do_resize (`bool`, *optional*, defaults to `True`):
973
+ Whether to downscale the image's (height, width) dimensions to multiples of `vae_scale_factor`.
974
+ vae_scale_factor (`int`, *optional*, defaults to `8`):
975
+ VAE scale factor. If `do_resize` is `True`, the image is automatically resized to multiples of this factor.
976
+ resample (`str`, *optional*, defaults to `lanczos`):
977
+ Resampling filter to use when resizing the image.
978
+ do_normalize (`bool`, *optional*, defaults to `True`):
979
+ Whether to normalize the image to [-1,1].
980
+ """
981
+
982
+ config_name = CONFIG_NAME
983
+
984
+ @register_to_config
985
+ def __init__(
986
+ self,
987
+ do_resize: bool = True,
988
+ vae_scale_factor: int = 8,
989
+ resample: str = "lanczos",
990
+ do_normalize: bool = True,
991
+ ):
992
+ super().__init__()
993
+
994
+ @staticmethod
995
+ def numpy_to_pil(images: np.ndarray) -> list[PIL.Image.Image]:
996
+ r"""
997
+ Convert a NumPy image or a batch of images to a list of PIL images.
998
+
999
+ Args:
1000
+ images (`np.ndarray`):
1001
+ The input NumPy array of images, which can be a single image or a batch.
1002
+
1003
+ Returns:
1004
+ `list[PIL.Image.Image]`:
1005
+ A list of PIL images converted from the input NumPy array.
1006
+ """
1007
+ if images.ndim == 3:
1008
+ images = images[None, ...]
1009
+ images = (images * 255).round().astype("uint8")
1010
+ if images.shape[-1] == 1:
1011
+ # special case for grayscale (single channel) images
1012
+ pil_images = [Image.fromarray(image.squeeze(), mode="L") for image in images]
1013
+ else:
1014
+ pil_images = [Image.fromarray(image[:, :, :3]) for image in images]
1015
+
1016
+ return pil_images
1017
+
1018
+ @staticmethod
1019
+ def depth_pil_to_numpy(images: list[PIL.Image.Image] | PIL.Image.Image) -> np.ndarray:
1020
+ r"""
1021
+ Convert a PIL image or a list of PIL images to NumPy arrays.
1022
+
1023
+ Args:
1024
+ images (`list[PIL.Image.Image, PIL.Image.Image]`):
1025
+ The input image or list of images to be converted.
1026
+
1027
+ Returns:
1028
+ `np.ndarray`:
1029
+ A NumPy array of the converted images.
1030
+ """
1031
+ if not isinstance(images, list):
1032
+ images = [images]
1033
+
1034
+ images = [np.array(image).astype(np.float32) / (2**16 - 1) for image in images]
1035
+ images = np.stack(images, axis=0)
1036
+ return images
1037
+
1038
+ @staticmethod
1039
+ def rgblike_to_depthmap(image: np.ndarray | torch.Tensor) -> np.ndarray | torch.Tensor:
1040
+ r"""
1041
+ Convert an RGB-like depth image to a depth map.
1042
+ """
1043
+ # 1. Cast the tensor to a larger integer type (e.g., int32)
1044
+ # to safely perform the multiplication by 256.
1045
+ # 2. Perform the 16-bit combination: High-byte * 256 + Low-byte.
1046
+ # 3. Cast the final result to the desired depth map type (uint16) if needed
1047
+ # before returning, though leaving it as int32/int64 is often safer
1048
+ # for return value from a library function.
1049
+
1050
+ if isinstance(image, torch.Tensor):
1051
+ # Cast to a safe dtype (e.g., int32 or int64) for the calculation
1052
+ original_dtype = image.dtype
1053
+ image_safe = image.to(torch.int32)
1054
+
1055
+ # Calculate the depth map
1056
+ depth_map = image_safe[:, :, 1] * 256 + image_safe[:, :, 2]
1057
+
1058
+ # You may want to cast the final result to uint16, but casting to a
1059
+ # larger int type (like int32) is sufficient to fix the overflow.
1060
+ # depth_map = depth_map.to(torch.uint16) # Uncomment if uint16 is strictly required
1061
+ return depth_map.to(original_dtype)
1062
+
1063
+ elif isinstance(image, np.ndarray):
1064
+ # NumPy equivalent: Cast to a safe dtype (e.g., np.int32)
1065
+ original_dtype = image.dtype
1066
+ image_safe = image.astype(np.int32)
1067
+
1068
+ # Calculate the depth map
1069
+ depth_map = image_safe[:, :, 1] * 256 + image_safe[:, :, 2]
1070
+
1071
+ # depth_map = depth_map.astype(np.uint16) # Uncomment if uint16 is strictly required
1072
+ return depth_map.astype(original_dtype)
1073
+ else:
1074
+ raise TypeError("Input image must be a torch.Tensor or np.ndarray")
1075
+
1076
+ def numpy_to_depth(self, images: np.ndarray) -> list[PIL.Image.Image]:
1077
+ r"""
1078
+ Convert a NumPy depth image or a batch of images to a list of PIL images.
1079
+
1080
+ Args:
1081
+ images (`np.ndarray`):
1082
+ The input NumPy array of depth images, which can be a single image or a batch.
1083
+
1084
+ Returns:
1085
+ `list[PIL.Image.Image]`:
1086
+ A list of PIL images converted from the input NumPy depth images.
1087
+ """
1088
+ if images.ndim == 3:
1089
+ images = images[None, ...]
1090
+ images_depth = images[:, :, :, 3:]
1091
+ if images.shape[-1] == 6:
1092
+ images_depth = (images_depth * 255).round().astype("uint8")
1093
+ pil_images = [
1094
+ Image.fromarray(self.rgblike_to_depthmap(image_depth), mode="I;16") for image_depth in images_depth
1095
+ ]
1096
+ elif images.shape[-1] == 4:
1097
+ images_depth = (images_depth * 65535.0).astype(np.uint16)
1098
+ pil_images = [Image.fromarray(image_depth, mode="I;16") for image_depth in images_depth]
1099
+ else:
1100
+ raise Exception("Not supported")
1101
+
1102
+ return pil_images
1103
+
1104
+ def postprocess(
1105
+ self,
1106
+ image: torch.Tensor,
1107
+ output_type: str = "pil",
1108
+ do_denormalize: list[bool] | None = None,
1109
+ ) -> PIL.Image.Image | np.ndarray | torch.Tensor:
1110
+ """
1111
+ Postprocess the image output from tensor to `output_type`.
1112
+
1113
+ Args:
1114
+ image (`torch.Tensor`):
1115
+ The image input, should be a pytorch tensor with shape `B x C x H x W`.
1116
+ output_type (`str`, *optional*, defaults to `pil`):
1117
+ The output type of the image, can be one of `pil`, `np`, `pt`, `latent`.
1118
+ do_denormalize (`list[bool]`, *optional*, defaults to `None`):
1119
+ Whether to denormalize the image to [0,1]. If `None`, will use the value of `do_normalize` in the
1120
+ `VaeImageProcessor` config.
1121
+
1122
+ Returns:
1123
+ `PIL.Image.Image`, `np.ndarray` or `torch.Tensor`:
1124
+ The postprocessed image.
1125
+ """
1126
+ if not isinstance(image, torch.Tensor):
1127
+ raise ValueError(
1128
+ f"Input for postprocessing is in incorrect format: {type(image)}. We only support pytorch tensor"
1129
+ )
1130
+ if output_type not in ["latent", "pt", "np", "pil"]:
1131
+ deprecation_message = (
1132
+ f"the output_type {output_type} is outdated and has been set to `np`. Please make sure to set it to one of these instead: "
1133
+ "`pil`, `np`, `pt`, `latent`"
1134
+ )
1135
+ deprecate("Unsupported output_type", "1.0.0", deprecation_message, standard_warn=False)
1136
+ output_type = "np"
1137
+
1138
+ image = self._denormalize_conditionally(image, do_denormalize)
1139
+
1140
+ image = self.pt_to_numpy(image)
1141
+
1142
+ if output_type == "np":
1143
+ if image.shape[-1] == 6:
1144
+ image_depth = np.stack([self.rgblike_to_depthmap(im[:, :, 3:]) for im in image], axis=0)
1145
+ else:
1146
+ image_depth = image[:, :, :, 3:]
1147
+ return image[:, :, :, :3], image_depth
1148
+
1149
+ if output_type == "pil":
1150
+ return self.numpy_to_pil(image), self.numpy_to_depth(image)
1151
+ else:
1152
+ raise Exception(f"This type {output_type} is not supported")
1153
+
1154
+ def preprocess(
1155
+ self,
1156
+ rgb: torch.Tensor | PIL.Image.Image | np.ndarray,
1157
+ depth: torch.Tensor | PIL.Image.Image | np.ndarray,
1158
+ height: int | None = None,
1159
+ width: int | None = None,
1160
+ target_res: int | None = None,
1161
+ ) -> torch.Tensor:
1162
+ r"""
1163
+ Preprocess the image input. Accepted formats are PIL images, NumPy arrays, or PyTorch tensors.
1164
+
1165
+ Args:
1166
+ rgb (`torch.Tensor | PIL.Image.Image | np.ndarray`):
1167
+ The RGB input image, which can be a single image or a batch.
1168
+ depth (`torch.Tensor | PIL.Image.Image | np.ndarray`):
1169
+ The depth input image, which can be a single image or a batch.
1170
+ height (`int | None`, *optional*, defaults to `None`):
1171
+ The desired height of the processed image. If `None`, defaults to the height of the input image.
1172
+ width (`int | None`, *optional*, defaults to `None`):
1173
+ The desired width of the processed image. If `None`, defaults to the width of the input image.
1174
+ target_res (`int | None`, *optional*, defaults to `None`):
1175
+ Target resolution for resizing the images. If specified, overrides height and width.
1176
+
1177
+ Returns:
1178
+ `tuple[torch.Tensor, torch.Tensor]`:
1179
+ A tuple containing the processed RGB and depth images as PyTorch tensors.
1180
+ """
1181
+ supported_formats = (PIL.Image.Image, np.ndarray, torch.Tensor)
1182
+
1183
+ # Expand the missing dimension for 3-dimensional pytorch tensor or numpy array that represents grayscale image
1184
+ if self.config.do_convert_grayscale and isinstance(rgb, (torch.Tensor, np.ndarray)) and rgb.ndim == 3:
1185
+ raise Exception("This is not yet supported")
1186
+
1187
+ if isinstance(rgb, supported_formats):
1188
+ rgb = [rgb]
1189
+ depth = [depth]
1190
+ elif not (isinstance(rgb, list) and all(isinstance(i, supported_formats) for i in rgb)):
1191
+ raise ValueError(
1192
+ f"Input is in incorrect format: {[type(i) for i in rgb]}. Currently, we only support {', '.join(supported_formats)}"
1193
+ )
1194
+
1195
+ if isinstance(rgb[0], PIL.Image.Image):
1196
+ if self.config.do_convert_rgb:
1197
+ raise Exception("This is not yet supported")
1198
+ # rgb = [self.convert_to_rgb(i) for i in rgb]
1199
+ # depth = [self.convert_to_depth(i) for i in depth] #TODO define convert_to_depth
1200
+ if self.config.do_resize or target_res:
1201
+ height, width = self.get_default_height_width(rgb[0], height, width) if not target_res else target_res
1202
+ rgb = [self.resize(i, height, width) for i in rgb]
1203
+ depth = [self.resize(i, height, width) for i in depth]
1204
+ rgb = self.pil_to_numpy(rgb) # to np
1205
+ rgb = self.numpy_to_pt(rgb) # to pt
1206
+
1207
+ depth = self.depth_pil_to_numpy(depth) # to np
1208
+ depth = self.numpy_to_pt(depth) # to pt
1209
+
1210
+ elif isinstance(rgb[0], np.ndarray):
1211
+ rgb = np.concatenate(rgb, axis=0) if rgb[0].ndim == 4 else np.stack(rgb, axis=0)
1212
+ rgb = self.numpy_to_pt(rgb)
1213
+ height, width = self.get_default_height_width(rgb, height, width)
1214
+ if self.config.do_resize:
1215
+ rgb = self.resize(rgb, height, width)
1216
+
1217
+ depth = np.concatenate(depth, axis=0) if rgb[0].ndim == 4 else np.stack(depth, axis=0)
1218
+ depth = self.numpy_to_pt(depth)
1219
+ height, width = self.get_default_height_width(depth, height, width)
1220
+ if self.config.do_resize:
1221
+ depth = self.resize(depth, height, width)
1222
+
1223
+ elif isinstance(rgb[0], torch.Tensor):
1224
+ raise Exception("This is not yet supported")
1225
+ # rgb = torch.cat(rgb, axis=0) if rgb[0].ndim == 4 else torch.stack(rgb, axis=0)
1226
+
1227
+ # if self.config.do_convert_grayscale and rgb.ndim == 3:
1228
+ # rgb = rgb.unsqueeze(1)
1229
+
1230
+ # channel = rgb.shape[1]
1231
+
1232
+ # height, width = self.get_default_height_width(rgb, height, width)
1233
+ # if self.config.do_resize:
1234
+ # rgb = self.resize(rgb, height, width)
1235
+
1236
+ # depth = torch.cat(depth, axis=0) if depth[0].ndim == 4 else torch.stack(depth, axis=0)
1237
+
1238
+ # if self.config.do_convert_grayscale and depth.ndim == 3:
1239
+ # depth = depth.unsqueeze(1)
1240
+
1241
+ # channel = depth.shape[1]
1242
+ # # don't need any preprocess if the image is latents
1243
+ # if depth == 4:
1244
+ # return rgb, depth
1245
+
1246
+ # height, width = self.get_default_height_width(depth, height, width)
1247
+ # if self.config.do_resize:
1248
+ # depth = self.resize(depth, height, width)
1249
+ # expected range [0,1], normalize to [-1,1]
1250
+ do_normalize = self.config.do_normalize
1251
+ if rgb.min() < 0 and do_normalize:
1252
+ warnings.warn(
1253
+ "Passing `image` as torch tensor with value range in [-1,1] is deprecated. The expected value range for image tensor is [0,1] "
1254
+ f"when passing as pytorch tensor or numpy Array. You passed `image` with value range [{rgb.min()},{rgb.max()}]",
1255
+ FutureWarning,
1256
+ )
1257
+ do_normalize = False
1258
+
1259
+ if do_normalize:
1260
+ rgb = self.normalize(rgb)
1261
+ depth = self.normalize(depth)
1262
+
1263
+ if self.config.do_binarize:
1264
+ rgb = self.binarize(rgb)
1265
+ depth = self.binarize(depth)
1266
+
1267
+ return rgb, depth
1268
+
1269
+
1270
+ class IPAdapterMaskProcessor(VaeImageProcessor):
1271
+ """
1272
+ Image processor for IP Adapter image masks.
1273
+
1274
+ Args:
1275
+ do_resize (`bool`, *optional*, defaults to `True`):
1276
+ Whether to downscale the image's (height, width) dimensions to multiples of `vae_scale_factor`.
1277
+ vae_scale_factor (`int`, *optional*, defaults to `8`):
1278
+ VAE scale factor. If `do_resize` is `True`, the image is automatically resized to multiples of this factor.
1279
+ resample (`str`, *optional*, defaults to `lanczos`):
1280
+ Resampling filter to use when resizing the image.
1281
+ do_normalize (`bool`, *optional*, defaults to `False`):
1282
+ Whether to normalize the image to [-1,1].
1283
+ do_binarize (`bool`, *optional*, defaults to `True`):
1284
+ Whether to binarize the image to 0/1.
1285
+ do_convert_grayscale (`bool`, *optional*, defaults to be `True`):
1286
+ Whether to convert the images to grayscale format.
1287
+
1288
+ """
1289
+
1290
+ config_name = CONFIG_NAME
1291
+
1292
+ @register_to_config
1293
+ def __init__(
1294
+ self,
1295
+ do_resize: bool = True,
1296
+ vae_scale_factor: int = 8,
1297
+ resample: str = "lanczos",
1298
+ do_normalize: bool = False,
1299
+ do_binarize: bool = True,
1300
+ do_convert_grayscale: bool = True,
1301
+ ):
1302
+ super().__init__(
1303
+ do_resize=do_resize,
1304
+ vae_scale_factor=vae_scale_factor,
1305
+ resample=resample,
1306
+ do_normalize=do_normalize,
1307
+ do_binarize=do_binarize,
1308
+ do_convert_grayscale=do_convert_grayscale,
1309
+ )
1310
+
1311
+ @staticmethod
1312
+ def downsample(mask: torch.Tensor, batch_size: int, num_queries: int, value_embed_dim: int):
1313
+ """
1314
+ Downsamples the provided mask tensor to match the expected dimensions for scaled dot-product attention. If the
1315
+ aspect ratio of the mask does not match the aspect ratio of the output image, a warning is issued.
1316
+
1317
+ Args:
1318
+ mask (`torch.Tensor`):
1319
+ The input mask tensor generated with `IPAdapterMaskProcessor.preprocess()`.
1320
+ batch_size (`int`):
1321
+ The batch size.
1322
+ num_queries (`int`):
1323
+ The number of queries.
1324
+ value_embed_dim (`int`):
1325
+ The dimensionality of the value embeddings.
1326
+
1327
+ Returns:
1328
+ `torch.Tensor`:
1329
+ The downsampled mask tensor.
1330
+
1331
+ """
1332
+ o_h = mask.shape[1]
1333
+ o_w = mask.shape[2]
1334
+ ratio = o_w / o_h
1335
+ mask_h = int(math.sqrt(num_queries / ratio))
1336
+ mask_h = int(mask_h) + int((num_queries % int(mask_h)) != 0)
1337
+ mask_w = num_queries // mask_h
1338
+
1339
+ mask_downsample = F.interpolate(mask.unsqueeze(0), size=(mask_h, mask_w), mode="bicubic").squeeze(0)
1340
+
1341
+ # Repeat batch_size times
1342
+ if mask_downsample.shape[0] < batch_size:
1343
+ mask_downsample = mask_downsample.repeat(batch_size, 1, 1)
1344
+
1345
+ mask_downsample = mask_downsample.view(mask_downsample.shape[0], -1)
1346
+
1347
+ downsampled_area = mask_h * mask_w
1348
+ # If the output image and the mask do not have the same aspect ratio, tensor shapes will not match
1349
+ # Pad tensor if downsampled_mask.shape[1] is smaller than num_queries
1350
+ if downsampled_area < num_queries:
1351
+ warnings.warn(
1352
+ "The aspect ratio of the mask does not match the aspect ratio of the output image. "
1353
+ "Please update your masks or adjust the output size for optimal performance.",
1354
+ UserWarning,
1355
+ )
1356
+ mask_downsample = F.pad(mask_downsample, (0, num_queries - mask_downsample.shape[1]), value=0.0)
1357
+ # Discard last embeddings if downsampled_mask.shape[1] is bigger than num_queries
1358
+ if downsampled_area > num_queries:
1359
+ warnings.warn(
1360
+ "The aspect ratio of the mask does not match the aspect ratio of the output image. "
1361
+ "Please update your masks or adjust the output size for optimal performance.",
1362
+ UserWarning,
1363
+ )
1364
+ mask_downsample = mask_downsample[:, :num_queries]
1365
+
1366
+ # Repeat last dimension to match SDPA output shape
1367
+ mask_downsample = mask_downsample.view(mask_downsample.shape[0], mask_downsample.shape[1], 1).repeat(
1368
+ 1, 1, value_embed_dim
1369
+ )
1370
+
1371
+ return mask_downsample
1372
+
1373
+
1374
+ class PixArtImageProcessor(VaeImageProcessor):
1375
+ """
1376
+ Image processor for PixArt image resize and crop.
1377
+
1378
+ Args:
1379
+ do_resize (`bool`, *optional*, defaults to `True`):
1380
+ Whether to downscale the image's (height, width) dimensions to multiples of `vae_scale_factor`. Can accept
1381
+ `height` and `width` arguments from [`image_processor.VaeImageProcessor.preprocess`] method.
1382
+ vae_scale_factor (`int`, *optional*, defaults to `8`):
1383
+ VAE scale factor. If `do_resize` is `True`, the image is automatically resized to multiples of this factor.
1384
+ resample (`str`, *optional*, defaults to `lanczos`):
1385
+ Resampling filter to use when resizing the image.
1386
+ do_normalize (`bool`, *optional*, defaults to `True`):
1387
+ Whether to normalize the image to [-1,1].
1388
+ do_binarize (`bool`, *optional*, defaults to `False`):
1389
+ Whether to binarize the image to 0/1.
1390
+ do_convert_rgb (`bool`, *optional*, defaults to be `False`):
1391
+ Whether to convert the images to RGB format.
1392
+ do_convert_grayscale (`bool`, *optional*, defaults to be `False`):
1393
+ Whether to convert the images to grayscale format.
1394
+ """
1395
+
1396
+ @register_to_config
1397
+ def __init__(
1398
+ self,
1399
+ do_resize: bool = True,
1400
+ vae_scale_factor: int = 8,
1401
+ resample: str = "lanczos",
1402
+ do_normalize: bool = True,
1403
+ do_binarize: bool = False,
1404
+ do_convert_grayscale: bool = False,
1405
+ ):
1406
+ super().__init__(
1407
+ do_resize=do_resize,
1408
+ vae_scale_factor=vae_scale_factor,
1409
+ resample=resample,
1410
+ do_normalize=do_normalize,
1411
+ do_binarize=do_binarize,
1412
+ do_convert_grayscale=do_convert_grayscale,
1413
+ )
1414
+
1415
+ @staticmethod
1416
+ def classify_height_width_bin(height: int, width: int, ratios: dict) -> tuple[int, int]:
1417
+ r"""
1418
+ Returns the binned height and width based on the aspect ratio.
1419
+
1420
+ Args:
1421
+ height (`int`): The height of the image.
1422
+ width (`int`): The width of the image.
1423
+ ratios (`dict`): A dictionary where keys are aspect ratios and values are tuples of (height, width).
1424
+
1425
+ Returns:
1426
+ `tuple[int, int]`: The closest binned height and width.
1427
+ """
1428
+ ar = float(height / width)
1429
+ closest_ratio = min(ratios.keys(), key=lambda ratio: abs(float(ratio) - ar))
1430
+ default_hw = ratios[closest_ratio]
1431
+ return int(default_hw[0]), int(default_hw[1])
1432
+
1433
+ @staticmethod
1434
+ def resize_and_crop_tensor(samples: torch.Tensor, new_width: int, new_height: int) -> torch.Tensor:
1435
+ r"""
1436
+ Resizes and crops a tensor of images to the specified dimensions.
1437
+
1438
+ Args:
1439
+ samples (`torch.Tensor`):
1440
+ A tensor of shape (N, C, H, W) where N is the batch size, C is the number of channels, H is the height,
1441
+ and W is the width.
1442
+ new_width (`int`): The desired width of the output images.
1443
+ new_height (`int`): The desired height of the output images.
1444
+
1445
+ Returns:
1446
+ `torch.Tensor`: A tensor containing the resized and cropped images.
1447
+ """
1448
+ orig_height, orig_width = samples.shape[2], samples.shape[3]
1449
+
1450
+ # Check if resizing is needed
1451
+ if orig_height != new_height or orig_width != new_width:
1452
+ ratio = max(new_height / orig_height, new_width / orig_width)
1453
+ resized_width = int(orig_width * ratio)
1454
+ resized_height = int(orig_height * ratio)
1455
+
1456
+ # Resize
1457
+ samples = F.interpolate(
1458
+ samples, size=(resized_height, resized_width), mode="bilinear", align_corners=False
1459
+ )
1460
+
1461
+ # Center Crop
1462
+ start_x = (resized_width - new_width) // 2
1463
+ end_x = start_x + new_width
1464
+ start_y = (resized_height - new_height) // 2
1465
+ end_y = start_y + new_height
1466
+ samples = samples[:, :, start_y:end_y, start_x:end_x]
1467
+
1468
+ return samples
diffusers/loaders/__init__.py ADDED
@@ -0,0 +1,159 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from typing import TYPE_CHECKING
2
+
3
+ from ..utils import DIFFUSERS_SLOW_IMPORT, _LazyModule, deprecate
4
+ from ..utils.import_utils import is_peft_available, is_torch_available, is_transformers_available
5
+
6
+
7
+ def text_encoder_lora_state_dict(text_encoder):
8
+ deprecate(
9
+ "text_encoder_load_state_dict in `models`",
10
+ "0.27.0",
11
+ "`text_encoder_lora_state_dict` is deprecated and will be removed in 0.27.0. Make sure to retrieve the weights using `get_peft_model`. See https://huggingface.co/docs/peft/v0.6.2/en/quicktour#peftmodel for more information.",
12
+ )
13
+ state_dict = {}
14
+
15
+ for name, module in text_encoder_attn_modules(text_encoder):
16
+ for k, v in module.q_proj.lora_linear_layer.state_dict().items():
17
+ state_dict[f"{name}.q_proj.lora_linear_layer.{k}"] = v
18
+
19
+ for k, v in module.k_proj.lora_linear_layer.state_dict().items():
20
+ state_dict[f"{name}.k_proj.lora_linear_layer.{k}"] = v
21
+
22
+ for k, v in module.v_proj.lora_linear_layer.state_dict().items():
23
+ state_dict[f"{name}.v_proj.lora_linear_layer.{k}"] = v
24
+
25
+ for k, v in module.out_proj.lora_linear_layer.state_dict().items():
26
+ state_dict[f"{name}.out_proj.lora_linear_layer.{k}"] = v
27
+
28
+ return state_dict
29
+
30
+
31
+ if is_transformers_available():
32
+
33
+ def text_encoder_attn_modules(text_encoder):
34
+ deprecate(
35
+ "text_encoder_attn_modules in `models`",
36
+ "0.27.0",
37
+ "`text_encoder_lora_state_dict` is deprecated and will be removed in 0.27.0. Make sure to retrieve the weights using `get_peft_model`. See https://huggingface.co/docs/peft/v0.6.2/en/quicktour#peftmodel for more information.",
38
+ )
39
+ from transformers import CLIPTextModel, CLIPTextModelWithProjection
40
+
41
+ attn_modules = []
42
+
43
+ if isinstance(text_encoder, (CLIPTextModel, CLIPTextModelWithProjection)):
44
+ for i, layer in enumerate(text_encoder.text_model.encoder.layers):
45
+ name = f"text_model.encoder.layers.{i}.self_attn"
46
+ mod = layer.self_attn
47
+ attn_modules.append((name, mod))
48
+ else:
49
+ raise ValueError(f"do not know how to get attention modules for: {text_encoder.__class__.__name__}")
50
+
51
+ return attn_modules
52
+
53
+
54
+ _import_structure = {}
55
+
56
+ if is_torch_available():
57
+ _import_structure["single_file_model"] = ["FromOriginalModelMixin"]
58
+ _import_structure["transformer_flux"] = ["FluxTransformer2DLoadersMixin"]
59
+ _import_structure["transformer_sd3"] = ["SD3Transformer2DLoadersMixin"]
60
+ _import_structure["unet"] = ["UNet2DConditionLoadersMixin"]
61
+ _import_structure["utils"] = ["AttnProcsLayers"]
62
+ if is_transformers_available():
63
+ _import_structure["single_file"] = ["FromSingleFileMixin"]
64
+ _import_structure["lora_pipeline"] = [
65
+ "AceStepLoraLoaderMixin",
66
+ "AmusedLoraLoaderMixin",
67
+ "AnimaLoraLoaderMixin",
68
+ "StableDiffusionLoraLoaderMixin",
69
+ "SD3LoraLoaderMixin",
70
+ "AuraFlowLoraLoaderMixin",
71
+ "StableDiffusionXLLoraLoaderMixin",
72
+ "LTX2LoraLoaderMixin",
73
+ "LTXVideoLoraLoaderMixin",
74
+ "LoraLoaderMixin",
75
+ "FluxLoraLoaderMixin",
76
+ "CogVideoXLoraLoaderMixin",
77
+ "CogView4LoraLoaderMixin",
78
+ "Mochi1LoraLoaderMixin",
79
+ "HunyuanVideoLoraLoaderMixin",
80
+ "SanaLoraLoaderMixin",
81
+ "Lumina2LoraLoaderMixin",
82
+ "WanLoraLoaderMixin",
83
+ "HeliosLoraLoaderMixin",
84
+ "KandinskyLoraLoaderMixin",
85
+ "HiDreamImageLoraLoaderMixin",
86
+ "SkyReelsV2LoraLoaderMixin",
87
+ "QwenImageLoraLoaderMixin",
88
+ "Krea2LoraLoaderMixin",
89
+ "ZImageLoraLoaderMixin",
90
+ "Flux2LoraLoaderMixin",
91
+ "Ideogram4LoraLoaderMixin",
92
+ "ErnieImageLoraLoaderMixin",
93
+ "CosmosLoraLoaderMixin",
94
+ ]
95
+ _import_structure["textual_inversion"] = ["TextualInversionLoaderMixin"]
96
+ _import_structure["ip_adapter"] = [
97
+ "IPAdapterMixin",
98
+ "FluxIPAdapterMixin",
99
+ "SD3IPAdapterMixin",
100
+ "ModularIPAdapterMixin",
101
+ ]
102
+
103
+ _import_structure["peft"] = ["PeftAdapterMixin"]
104
+
105
+
106
+ if TYPE_CHECKING or DIFFUSERS_SLOW_IMPORT:
107
+ if is_torch_available():
108
+ from .single_file_model import FromOriginalModelMixin
109
+ from .transformer_flux import FluxTransformer2DLoadersMixin
110
+ from .transformer_sd3 import SD3Transformer2DLoadersMixin
111
+ from .unet import UNet2DConditionLoadersMixin
112
+ from .utils import AttnProcsLayers
113
+
114
+ if is_transformers_available():
115
+ from .ip_adapter import (
116
+ FluxIPAdapterMixin,
117
+ IPAdapterMixin,
118
+ ModularIPAdapterMixin,
119
+ SD3IPAdapterMixin,
120
+ )
121
+ from .lora_pipeline import (
122
+ AceStepLoraLoaderMixin,
123
+ AmusedLoraLoaderMixin,
124
+ AnimaLoraLoaderMixin,
125
+ AuraFlowLoraLoaderMixin,
126
+ CogVideoXLoraLoaderMixin,
127
+ CogView4LoraLoaderMixin,
128
+ CosmosLoraLoaderMixin,
129
+ ErnieImageLoraLoaderMixin,
130
+ Flux2LoraLoaderMixin,
131
+ FluxLoraLoaderMixin,
132
+ HeliosLoraLoaderMixin,
133
+ HiDreamImageLoraLoaderMixin,
134
+ HunyuanVideoLoraLoaderMixin,
135
+ Ideogram4LoraLoaderMixin,
136
+ KandinskyLoraLoaderMixin,
137
+ Krea2LoraLoaderMixin,
138
+ LoraLoaderMixin,
139
+ LTX2LoraLoaderMixin,
140
+ LTXVideoLoraLoaderMixin,
141
+ Lumina2LoraLoaderMixin,
142
+ Mochi1LoraLoaderMixin,
143
+ QwenImageLoraLoaderMixin,
144
+ SanaLoraLoaderMixin,
145
+ SD3LoraLoaderMixin,
146
+ SkyReelsV2LoraLoaderMixin,
147
+ StableDiffusionLoraLoaderMixin,
148
+ StableDiffusionXLLoraLoaderMixin,
149
+ WanLoraLoaderMixin,
150
+ ZImageLoraLoaderMixin,
151
+ )
152
+ from .single_file import FromSingleFileMixin
153
+ from .textual_inversion import TextualInversionLoaderMixin
154
+
155
+ from .peft import PeftAdapterMixin
156
+ else:
157
+ import sys
158
+
159
+ sys.modules[__name__] = _LazyModule(__name__, globals()["__file__"], _import_structure, module_spec=__spec__)