multimodalart HF Staff commited on
Commit
9da256f
·
verified ·
1 Parent(s): ef0bfcc

Migrate to canonical diffusers PR 14371 + public MiniMaxAI/MiniMax-H3 weights, add prompt upsampling toggle

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 +31 -20
  2. app.py +52 -23
  3. diffusers/__init__.py +0 -1750
  4. diffusers/callbacks.py +0 -244
  5. diffusers/commands/__init__.py +0 -27
  6. diffusers/commands/custom_blocks.py +0 -140
  7. diffusers/commands/diffusers_cli.py +0 -69
  8. diffusers/commands/env.py +0 -185
  9. diffusers/commands/fp16_safetensors.py +0 -144
  10. diffusers/commands/run.py +0 -1227
  11. diffusers/commands/schema.py +0 -287
  12. diffusers/commands/skills.py +0 -344
  13. diffusers/configuration_utils.py +0 -752
  14. diffusers/dependency_versions_check.py +0 -34
  15. diffusers/dependency_versions_table.py +0 -57
  16. diffusers/experimental/README.md +0 -5
  17. diffusers/experimental/__init__.py +0 -1
  18. diffusers/experimental/rl/__init__.py +0 -1
  19. diffusers/experimental/rl/value_guided_sampling.py +0 -153
  20. diffusers/guiders/__init__.py +0 -31
  21. diffusers/guiders/adaptive_projected_guidance.py +0 -253
  22. diffusers/guiders/adaptive_projected_guidance_mix.py +0 -297
  23. diffusers/guiders/auto_guidance.py +0 -198
  24. diffusers/guiders/classifier_free_guidance.py +0 -156
  25. diffusers/guiders/classifier_free_zero_star_guidance.py +0 -164
  26. diffusers/guiders/frequency_decoupled_guidance.py +0 -335
  27. diffusers/guiders/guider_utils.py +0 -396
  28. diffusers/guiders/magnitude_aware_guidance.py +0 -159
  29. diffusers/guiders/perturbed_attention_guidance.py +0 -289
  30. diffusers/guiders/skip_layer_guidance.py +0 -280
  31. diffusers/guiders/smoothed_energy_guidance.py +0 -269
  32. diffusers/guiders/tangential_classifier_free_guidance.py +0 -151
  33. diffusers/hooks/__init__.py +0 -30
  34. diffusers/hooks/_common.py +0 -61
  35. diffusers/hooks/_helpers.py +0 -401
  36. diffusers/hooks/context_parallel.py +0 -382
  37. diffusers/hooks/faster_cache.py +0 -654
  38. diffusers/hooks/first_block_cache.py +0 -258
  39. diffusers/hooks/group_offloading.py +0 -1056
  40. diffusers/hooks/hooks.py +0 -312
  41. diffusers/hooks/layer_skip.py +0 -263
  42. diffusers/hooks/layerwise_casting.py +0 -240
  43. diffusers/hooks/mag_cache.py +0 -468
  44. diffusers/hooks/pyramid_attention_broadcast.py +0 -314
  45. diffusers/hooks/smoothed_energy_guidance_utils.py +0 -166
  46. diffusers/hooks/taylorseer_cache.py +0 -345
  47. diffusers/hooks/text_kv_cache.py +0 -173
  48. diffusers/hooks/utils.py +0 -43
  49. diffusers/image_processor.py +0 -1468
  50. diffusers/loaders/__init__.py +0 -159
README.md CHANGED
@@ -17,20 +17,19 @@ Joint video **and** soundtrack out of a single denoising pass, at **bfloat16 wit
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 + 10.43 GiB float32 |
35
 
36
  Besides the quality argument, unquantized weights are the ones AoTI can export; an NVFP4 checkpoint cannot be
@@ -59,19 +58,22 @@ Space pays, where the whole cost is the traffic auto-offload has to move.
59
 
60
  ## How the split is expressed
61
 
62
- `MiniMaxH3Blocks` is a `SequentialPipelineBlocks` of eight steps:
 
63
 
64
  ```
65
- setup -> text_encoder -> vae_encoder -> prepare_layout -> prepare_latents -> set_timesteps -> denoise -> decode
66
  ```
67
 
 
 
68
  `h3_split_blocks.py` subclasses it with the `text_encoder` step removed. Dropping the step drops the three
69
  components it declares, so `load_components` resolves `transformer` / `vae` / `audio_vae` / the two schedulers out of
70
  the shared `modular_model_index.json` and never fetches the conditioner — and `prompt_embeds` and `text_token_tags`
71
  become ordinary required inputs of the pipeline call:
72
 
73
  ```py
74
- pipe = MiniMaxH3GeneratorBlocks().init_pipeline("diffusers-internal-dev/MiniMax-H3")
75
  pipe.load_components(dtype=torch.bfloat16)
76
  state = pipe(prompt_embeds=..., text_token_tags=..., height=768, width=1344, num_frames=124, num_inference_steps=30)
77
  ```
@@ -80,10 +82,12 @@ The wire format is exactly those two tensors — `(1, num_text_tokens, 5120)` bf
80
  carried as one safetensors file with the resolved `height` / `width` / `num_frames` in its metadata header. A
81
  text-only request is 246 KB of it; one 768x1344 keyframe adds 1016 vision rows and takes it to 10.7 MB.
82
 
83
- The `setup` step runs on **both** halves. It owns no component (PIL and arithmetic) and it resolves the canvas, the
84
- `17 * n + 5` frame count and the keyframes placed onto that canvas — which the conditioner needs to build its vision
85
- blocks and this Space needs to encode with the video VAE. It is deterministic, and the conditioner returns the plan
86
- it resolved so this Space pins the same canvas rather than re-deriving it.
 
 
87
 
88
  ## Nothing is paid for with GPU time
89
 
@@ -147,18 +151,25 @@ one-time `PIPE.to("cuda")` is inside the first row's 339 s and does not reappear
147
 
148
  | Variable | Default | Meaning |
149
  |---|---|---|
150
- | `H3_CONDITIONER` | `diffusers-internal-dev/minimax-h3-conditioner` | The Space this one asks for embeddings. |
 
151
  | `H3_PLACEMENT` | `lazy` | `lazy` moves all 72.16 GiB onto the card on the first GPU call and leaves it there; `offload` hands placement to `ComponentsManager.enable_auto_cpu_offload` instead. |
152
  | `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. |
153
  | `H3_GPU_DURATION` | `900` | Seconds per request; the pool applies a 1.5 duration factor. |
154
  | `H3_GPU_SIZE` | `xlarge` | ZeroGPU allocation size. `large` does not fit. |
155
 
156
- ## Required secret
157
 
158
- `HF_TOKEN` `diffusers-internal-dev/MiniMax-H3` is private, and so is the conditioner Space this one calls.
 
159
 
160
  ## Where diffusers comes from
161
 
162
- MiniMax-H3 is modular-only and not in a released `diffusers`, so the integration branch's `src/diffusers` tree is
163
- vendored here as a top-level `diffusers/` package; the working directory comes first on `sys.path`, so there is no
164
- install step. `requirements.txt` only carries what that tree imports.
 
 
 
 
 
 
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
+ [`qwen3vl-conditioner`](https://huggingface.co/spaces/multimodalart/qwen3vl-conditioner), which this
21
+ Space calls over the gradio API for every request. The weights are the public
22
+ [`MiniMaxAI/MiniMax-H3`](https://huggingface.co/MiniMaxAI/MiniMax-H3) diffusers checkpoint.
23
 
24
  ## Why split
25
 
26
  MiniMax-H3 is 195.9 GiB in bfloat16 and a ZeroGPU Space is evicted at **150 GB of storage**. An unquantized single
27
+ Space is therefore impossible, which is why quantized demos of it run NVFP4 or float8 weights. Cut the
28
+ `MiniMaxH3Blocks` sequence at its `text_encoder` step and both halves fit unquantized:
 
 
29
 
30
  | Space | Subfolders | Download | Resident |
31
  |---|---|---|---|
32
+ | [`qwen3vl-conditioner`](https://huggingface.co/spaces/multimodalart/qwen3vl-conditioner) | `text_encoder/` + `tokenizer/` + `processor/` | 66.7 GB | 62.15 GiB bf16 |
33
  | this one | `transformer/` + `vae/` + `audio_vae/` | 77.3 GB | 61.73 GiB bf16 + 10.43 GiB float32 |
34
 
35
  Besides the quality argument, unquantized weights are the ones AoTI can export; an NVFP4 checkpoint cannot be
 
58
 
59
  ## How the split is expressed
60
 
61
+ `MiniMaxH3Blocks` is a `SequentialPipelineBlocks` whose branches are picked per request — and per `workflow=` — from
62
+ the inputs:
63
 
64
  ```
65
+ before_encode -> text_encoder -> vae_encoder -> denoise -> after_denoise -> decode
66
  ```
67
 
68
+ where `denoise` is itself `prepare_layout -> prepare_latents -> set_timesteps -> denoise`.
69
+
70
  `h3_split_blocks.py` subclasses it with the `text_encoder` step removed. Dropping the step drops the three
71
  components it declares, so `load_components` resolves `transformer` / `vae` / `audio_vae` / the two schedulers out of
72
  the shared `modular_model_index.json` and never fetches the conditioner — and `prompt_embeds` and `text_token_tags`
73
  become ordinary required inputs of the pipeline call:
74
 
75
  ```py
76
+ pipe = MiniMaxH3GeneratorBlocks().init_pipeline("MiniMaxAI/MiniMax-H3")
77
  pipe.load_components(dtype=torch.bfloat16)
78
  state = pipe(prompt_embeds=..., text_token_tags=..., height=768, width=1344, num_frames=124, num_inference_steps=30)
79
  ```
 
82
  carried as one safetensors file with the resolved `height` / `width` / `num_frames` in its metadata header. A
83
  text-only request is 246 KB of it; one 768x1344 keyframe adds 1016 vision rows and takes it to 10.7 MB.
84
 
85
+ The keyframe `resize` step runs on **both** halves. It owns no pretrained component (PIL and arithmetic) and it puts
86
+ the keyframes onto the target canvas — which the conditioner needs to build its vision blocks and this Space needs to
87
+ encode with the video VAE. It is deterministic, and the conditioner returns the plan it resolved so this Space pins
88
+ the same canvas rather than re-deriving it. Two things that step no longer does, and that both halves therefore do
89
+ themselves: EXIF-transposing a keyframe into upright RGB, and snapping `num_frames` to `17 * n + 5` — the frame count
90
+ is resolved by the layout step, which lives on this side of the cut.
91
 
92
  ## Nothing is paid for with GPU time
93
 
 
151
 
152
  | Variable | Default | Meaning |
153
  |---|---|---|
154
+ | `H3_MODEL_REPO` | `MiniMaxAI/MiniMax-H3` | The diffusers-layout checkpoint. Public. |
155
+ | `H3_CONDITIONER` | `multimodalart/qwen3vl-conditioner` | The Space this one asks for embeddings. |
156
  | `H3_PLACEMENT` | `lazy` | `lazy` moves all 72.16 GiB onto the card on the first GPU call and leaves it there; `offload` hands placement to `ComponentsManager.enable_auto_cpu_offload` instead. |
157
  | `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. |
158
  | `H3_GPU_DURATION` | `900` | Seconds per request; the pool applies a 1.5 duration factor. |
159
  | `H3_GPU_SIZE` | `xlarge` | ZeroGPU allocation size. `large` does not fit. |
160
 
161
+ ## Secrets
162
 
163
+ The weights and the conditioner Space are both public, so neither needs a token. `HF_TOKEN` is still read by
164
+ `h3_aoti`, whose compiled-package repository is private — without it, set `H3_AOTI=0`.
165
 
166
  ## Where diffusers comes from
167
 
168
+ MiniMax-H3 is modular-only and not in a released `diffusers`, so `requirements.txt` installs it from the canonical
169
+ pull request, [huggingface/diffusers#14371](https://github.com/huggingface/diffusers/pull/14371), pinned to the
170
+ **commit** `665f5782` (`refs/pull/14371/head` at deploy time) rather than to the moving `minimax-h3-refactor` branch.
171
+
172
+ That PR is a WIP: it needs **re-pinning whenever it updates**, and `h3_split_blocks.py` — which subclasses its block
173
+ classes to cut the pipeline in two — has to be re-checked against the new head at the same time. The PR refactored
174
+ the blocks into one workflow-selected pipeline, so block names and the shape of the split are exactly what a new head
175
+ is liable to move.
app.py CHANGED
@@ -13,7 +13,7 @@ import traceback
13
  import spaces
14
  import gradio as gr
15
 
16
- MODEL_REPO = os.environ.get("H3_MODEL_REPO", "diffusers-internal-dev/MiniMax-H3")
17
  CONDITIONER_SPACE = os.environ.get("H3_CONDITIONER", "multimodalart/qwen3vl-conditioner")
18
  # `lazy` moves all 72.16 GiB onto the card on the first GPU call and leaves it there; `offload` hands placement to
19
  # `ComponentsManager.enable_auto_cpu_offload` instead. Neither puts anything on the card at *startup*, which is
@@ -85,8 +85,9 @@ def load_models() -> str | None:
85
  """Load the denoising half. At **startup**, but *not* onto the card.
86
 
87
  `MiniMaxH3GeneratorBlocks` declares `transformer`, `vae`, `audio_vae`, `scheduler`, `audio_scheduler` and
88
- `video_processor`, so `load_components` fetches exactly those subfolders out of the shared
89
- `modular_model_index.json` `text_encoder/` and `transformer_ref/` are never touched.
 
90
 
91
  Both autoencoders carry `_keep_in_fp32_modules` over every module, so the `dtype` below is refused for them and
92
  they stay float32: a bfloat16 audio VAE decodes the soundtrack roughly 20 dB too quiet.
@@ -105,11 +106,6 @@ def load_models() -> str | None:
105
  if PIPE is not None or LOAD_ERROR is not None:
106
  return LOAD_ERROR
107
 
108
- token = os.environ.get("HF_TOKEN")
109
- if not token:
110
- LOAD_ERROR = f"**`HF_TOKEN` secret is missing** and `{MODEL_REPO}` is private. Add it and restart."
111
- return LOAD_ERROR
112
-
113
  started = time.time()
114
  try:
115
  import torch
@@ -121,7 +117,9 @@ def load_models() -> str | None:
121
  blocks = MiniMaxH3GeneratorBlocks()
122
  print(f"[gen] loading {[c.name for c in blocks.expected_components]} from {MODEL_REPO} ...", flush=True)
123
  pipe = blocks.init_pipeline(MODEL_REPO, components_manager=manager, collection="h3")
124
- pipe.load_components(dtype=torch.bfloat16, token=token)
 
 
125
  pipe.transformer.set_attention_backend(ATTENTION)
126
 
127
  # Still startup, still free: an AoTI package carries no weights and opens its compiled archive lazily inside
@@ -182,8 +180,14 @@ def conditioner():
182
  return CLIENT
183
 
184
 
185
- def encode_remote(prompt, image_path, last_image_path, canvas, num_frames):
186
- """Ask the conditioner Space for `prompt_embeds` + `text_token_tags`. Off this Space's GPU time entirely."""
 
 
 
 
 
 
187
  from gradio_client import handle_file
188
  from safetensors import safe_open
189
 
@@ -193,6 +197,7 @@ def encode_remote(prompt, image_path, last_image_path, canvas, num_frames):
193
  last_image_path=handle_file(last_image_path) if last_image_path else None,
194
  canvas=canvas,
195
  num_frames=num_frames,
 
196
  api_name="/encode",
197
  )
198
  with safe_open(path, framework="pt") as handle:
@@ -250,7 +255,8 @@ def _generate(prompt_embeds, text_token_tags, image, last_image, height, width,
250
  return state.get("videos")[0], state.get("audio")[0].cpu(), state.get("sampling_rate")
251
 
252
 
253
- def generate(prompt, image_path=None, last_image_path=None, canvas=DEFAULT_CANVAS, duration=5, steps=28, seed=42, progress=gr.Progress(track_tqdm=True)):
 
254
  if LOAD_ERROR:
255
  raise gr.Error(LOAD_ERROR)
256
  if PIPE is None:
@@ -258,27 +264,34 @@ def generate(prompt, image_path=None, last_image_path=None, canvas=DEFAULT_CANVA
258
  if not prompt or not prompt.strip():
259
  raise gr.Error("MiniMax-H3 always takes a prompt, keyframes or not.")
260
 
261
- from PIL import Image
262
 
263
  from diffusers.utils import encode_video
264
 
265
  num_frames = snap_frames(duration)
266
 
267
- progress(0.0, desc=f"Conditioning on {CONDITIONER_SPACE} ...")
268
  conditioned = time.time()
269
  prompt_embeds, text_token_tags, metadata, plan = encode_remote(
270
- prompt, image_path, last_image_path, canvas, num_frames
271
  )
272
  condition_seconds = time.time() - conditioned
273
  height, width, num_frames = (int(metadata[key]) for key in ("height", "width", "num_frames"))
 
 
 
 
 
 
 
274
 
275
  progress(0.1, desc=f"Denoising {steps} steps at {width}x{height}, {num_frames} frames ...")
276
  started = time.time()
277
  frames, audio, sampling_rate = _generate(
278
  prompt_embeds,
279
  text_token_tags,
280
- Image.open(image_path) if image_path else None,
281
- Image.open(last_image_path) if last_image_path else None,
282
  height,
283
  width,
284
  num_frames,
@@ -294,11 +307,12 @@ def generate(prompt, image_path=None, last_image_path=None, canvas=DEFAULT_CANVA
294
 
295
  report = (
296
  f"`{width}x{height}`, {num_frames} frames ({num_frames / FPS:.3f} s), {int(steps)} steps · "
297
- f"conditioner {condition_seconds:.0f}s ({plan['num_text_tokens']} tokens) · "
 
298
  f"denoise + decode {generate_seconds:.0f}s ({generate_seconds / int(steps):.1f} s/step) · seed {int(seed)}"
299
  )
300
  print(f"[gen] {report}", flush=True)
301
- return path, report
302
 
303
 
304
 
@@ -378,10 +392,22 @@ with gr.Blocks(title="MiniMax-H3") as demo:
378
  duration = gr.Slider(label="Duration (s)", minimum=2, maximum=MAX_UI_DURATION, step=1, value=5)
379
  steps = gr.Slider(label="Steps", minimum=10, maximum=40, step=1, value=28)
380
  seed = gr.Number(label="Seed", value=42, precision=0)
381
-
 
 
 
 
 
382
  with gr.Column():
383
  video = gr.Video(label="Video + soundtrack")
384
  report = gr.Markdown(visible=False)
 
 
 
 
 
 
 
385
 
386
  image.upload(_fit_keyframe, [image, canvas], [image, canvas])
387
 
@@ -394,16 +420,19 @@ with gr.Blocks(title="MiniMax-H3") as demo:
394
  ["A slow seamless camera move from the first view to the last", "examples/first.png", "examples/last.png", "1344x768 · 16:9 full"],
395
  ],
396
  inputs=[prompt, image, last_image, canvas],
397
- outputs=[video, report],
398
  fn=generate,
399
  cache_examples=True,
400
  cache_mode="lazy",
401
  )
402
 
 
 
 
403
  run.click(
404
  generate,
405
- [prompt, image, last_image, canvas, duration, steps, seed],
406
- [video, report],
407
  api_name="generate",
408
  )
409
 
 
13
  import spaces
14
  import gradio as gr
15
 
16
+ MODEL_REPO = os.environ.get("H3_MODEL_REPO", "MiniMaxAI/MiniMax-H3")
17
  CONDITIONER_SPACE = os.environ.get("H3_CONDITIONER", "multimodalart/qwen3vl-conditioner")
18
  # `lazy` moves all 72.16 GiB onto the card on the first GPU call and leaves it there; `offload` hands placement to
19
  # `ComponentsManager.enable_auto_cpu_offload` instead. Neither puts anything on the card at *startup*, which is
 
85
  """Load the denoising half. At **startup**, but *not* onto the card.
86
 
87
  `MiniMaxH3GeneratorBlocks` declares `transformer`, `vae`, `audio_vae`, `scheduler`, `audio_scheduler` and
88
+ `video_processor` as its pretrained components (plus an `image_processor` built from config), so
89
+ `load_components` fetches exactly those subfolders out of the shared `modular_model_index.json` —
90
+ `text_encoder/` and `transformer_ref/` are never touched.
91
 
92
  Both autoencoders carry `_keep_in_fp32_modules` over every module, so the `dtype` below is refused for them and
93
  they stay float32: a bfloat16 audio VAE decodes the soundtrack roughly 20 dB too quiet.
 
106
  if PIPE is not None or LOAD_ERROR is not None:
107
  return LOAD_ERROR
108
 
 
 
 
 
 
109
  started = time.time()
110
  try:
111
  import torch
 
117
  blocks = MiniMaxH3GeneratorBlocks()
118
  print(f"[gen] loading {[c.name for c in blocks.expected_components]} from {MODEL_REPO} ...", flush=True)
119
  pipe = blocks.init_pipeline(MODEL_REPO, components_manager=manager, collection="h3")
120
+ # `MiniMaxAI/MiniMax-H3` is public, so the weights are fetched without a token. `HF_TOKEN` is still read by
121
+ # `h3_aoti`, whose compiled-package repository is not.
122
+ pipe.load_components(dtype=torch.bfloat16)
123
  pipe.transformer.set_attention_backend(ATTENTION)
124
 
125
  # Still startup, still free: an AoTI package carries no weights and opens its compiled archive lazily inside
 
180
  return CLIENT
181
 
182
 
183
+ def encode_remote(prompt, image_path, last_image_path, canvas, num_frames, rewrite_prompt=False):
184
+ """Ask the conditioner Space for `prompt_embeds` + `text_token_tags`. Off this Space's GPU time entirely.
185
+
186
+ `rewrite_prompt` is the conditioner's prompt upsampling: it rewrites the request into MiniMax-H3's trained format
187
+ with its own Qwen3-VL and encodes *that*, handing the rewrite back under the plan's `refined_prompt`. It runs on
188
+ the conditioner's GPU booking, and this whole call happens before `_generate` books a card here, so it costs this
189
+ Space's `get_duration` nothing.
190
+ """
191
  from gradio_client import handle_file
192
  from safetensors import safe_open
193
 
 
197
  last_image_path=handle_file(last_image_path) if last_image_path else None,
198
  canvas=canvas,
199
  num_frames=num_frames,
200
+ rewrite_prompt=bool(rewrite_prompt),
201
  api_name="/encode",
202
  )
203
  with safe_open(path, framework="pt") as handle:
 
255
  return state.get("videos")[0], state.get("audio")[0].cpu(), state.get("sampling_rate")
256
 
257
 
258
+ def generate(prompt, image_path=None, last_image_path=None, canvas=DEFAULT_CANVAS, duration=5, steps=28, seed=42, upsample=False, progress=gr.Progress(track_tqdm=True)):
259
+ """One request. `upsample` is appended last and defaults off, so an existing API client is untouched by it."""
260
  if LOAD_ERROR:
261
  raise gr.Error(LOAD_ERROR)
262
  if PIPE is None:
 
264
  if not prompt or not prompt.strip():
265
  raise gr.Error("MiniMax-H3 always takes a prompt, keyframes or not.")
266
 
267
+ from PIL import Image, ImageOps
268
 
269
  from diffusers.utils import encode_video
270
 
271
  num_frames = snap_frames(duration)
272
 
273
+ progress(0.0, desc=f"Upsampling the prompt on {CONDITIONER_SPACE} ..." if upsample else f"Conditioning on {CONDITIONER_SPACE} ...")
274
  conditioned = time.time()
275
  prompt_embeds, text_token_tags, metadata, plan = encode_remote(
276
+ prompt, image_path, last_image_path, canvas, num_frames, rewrite_prompt=upsample
277
  )
278
  condition_seconds = time.time() - conditioned
279
  height, width, num_frames = (int(metadata[key]) for key in ("height", "width", "num_frames"))
280
+ refined = plan.get("refined_prompt") or ""
281
+
282
+ # EXIF-transposed and in RGB before the blocks see them, which the resize step that replaced the old setup step no
283
+ # longer does itself. Both halves have to prepare a keyframe the same way or the conditioning latents encoded here
284
+ # would not be of the image the conditioner looked at.
285
+ def keyframe(path):
286
+ return ImageOps.exif_transpose(Image.open(path)).convert("RGB") if path else None
287
 
288
  progress(0.1, desc=f"Denoising {steps} steps at {width}x{height}, {num_frames} frames ...")
289
  started = time.time()
290
  frames, audio, sampling_rate = _generate(
291
  prompt_embeds,
292
  text_token_tags,
293
+ keyframe(image_path),
294
+ keyframe(last_image_path),
295
  height,
296
  width,
297
  num_frames,
 
307
 
308
  report = (
309
  f"`{width}x{height}`, {num_frames} frames ({num_frames / FPS:.3f} s), {int(steps)} steps · "
310
+ f"conditioner {condition_seconds:.0f}s ({plan['num_text_tokens']} tokens"
311
+ f"{', upsampled' if refined else ''}) · "
312
  f"denoise + decode {generate_seconds:.0f}s ({generate_seconds / int(steps):.1f} s/step) · seed {int(seed)}"
313
  )
314
  print(f"[gen] {report}", flush=True)
315
+ return path, report, refined
316
 
317
 
318
 
 
392
  duration = gr.Slider(label="Duration (s)", minimum=2, maximum=MAX_UI_DURATION, step=1, value=5)
393
  steps = gr.Slider(label="Steps", minimum=10, maximum=40, step=1, value=28)
394
  seed = gr.Number(label="Seed", value=42, precision=0)
395
+ upsample = gr.Checkbox(
396
+ label="Upsample prompt",
397
+ value=False,
398
+ info="Rewrites the prompt into the model's trained format with the conditioner's Qwen3-VL before encoding.",
399
+ )
400
+
401
  with gr.Column():
402
  video = gr.Video(label="Video + soundtrack")
403
  report = gr.Markdown(visible=False)
404
+ with gr.Accordion("Upsampled prompt", open=False):
405
+ upsampled = gr.Textbox(
406
+ show_label=False,
407
+ lines=8,
408
+ interactive=False,
409
+ placeholder="Turn on “Upsample prompt” to see the rewrite that was encoded.",
410
+ )
411
 
412
  image.upload(_fit_keyframe, [image, canvas], [image, canvas])
413
 
 
420
  ["A slow seamless camera move from the first view to the last", "examples/first.png", "examples/last.png", "1344x768 · 16:9 full"],
421
  ],
422
  inputs=[prompt, image, last_image, canvas],
423
+ outputs=[video, report, upsampled],
424
  fn=generate,
425
  cache_examples=True,
426
  cache_mode="lazy",
427
  )
428
 
429
+ # `upsample` is appended *after* every input that was already here and every existing input keeps its position, so
430
+ # a positional API client that predates it keeps working and simply takes the default. Same on the way out: the
431
+ # video and the report stay first and the upsampled prompt is appended last.
432
  run.click(
433
  generate,
434
+ [prompt, image, last_image, canvas, duration, steps, seed, upsample],
435
+ [video, report, upsampled],
436
  api_name="generate",
437
  )
438
 
diffusers/__init__.py DELETED
@@ -1,1750 +0,0 @@
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 DELETED
@@ -1,244 +0,0 @@
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 DELETED
@@ -1,27 +0,0 @@
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 DELETED
@@ -1,140 +0,0 @@
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 DELETED
@@ -1,69 +0,0 @@
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 DELETED
@@ -1,185 +0,0 @@
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 DELETED
@@ -1,144 +0,0 @@
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 DELETED
@@ -1,1227 +0,0 @@
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 DELETED
@@ -1,287 +0,0 @@
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 DELETED
@@ -1,344 +0,0 @@
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 DELETED
@@ -1,752 +0,0 @@
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 DELETED
@@ -1,34 +0,0 @@
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 DELETED
@@ -1,57 +0,0 @@
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 DELETED
@@ -1,5 +0,0 @@
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 DELETED
@@ -1 +0,0 @@
1
- from .rl import ValueGuidedRLPipeline
 
 
diffusers/experimental/rl/__init__.py DELETED
@@ -1 +0,0 @@
1
- from .value_guided_sampling import ValueGuidedRLPipeline
 
 
diffusers/experimental/rl/value_guided_sampling.py DELETED
@@ -1,153 +0,0 @@
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 DELETED
@@ -1,31 +0,0 @@
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 DELETED
@@ -1,253 +0,0 @@
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 DELETED
@@ -1,297 +0,0 @@
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 DELETED
@@ -1,198 +0,0 @@
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 DELETED
@@ -1,156 +0,0 @@
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 DELETED
@@ -1,164 +0,0 @@
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 DELETED
@@ -1,335 +0,0 @@
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 DELETED
@@ -1,396 +0,0 @@
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 DELETED
@@ -1,159 +0,0 @@
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 DELETED
@@ -1,289 +0,0 @@
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 DELETED
@@ -1,280 +0,0 @@
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 DELETED
@@ -1,269 +0,0 @@
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 DELETED
@@ -1,151 +0,0 @@
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 DELETED
@@ -1,30 +0,0 @@
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 DELETED
@@ -1,61 +0,0 @@
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 DELETED
@@ -1,401 +0,0 @@
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 DELETED
@@ -1,382 +0,0 @@
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 DELETED
@@ -1,654 +0,0 @@
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 DELETED
@@ -1,258 +0,0 @@
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 DELETED
@@ -1,1056 +0,0 @@
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 DELETED
@@ -1,312 +0,0 @@
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 DELETED
@@ -1,263 +0,0 @@
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 DELETED
@@ -1,240 +0,0 @@
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 DELETED
@@ -1,468 +0,0 @@
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 DELETED
@@ -1,314 +0,0 @@
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 DELETED
@@ -1,166 +0,0 @@
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 DELETED
@@ -1,345 +0,0 @@
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 DELETED
@@ -1,173 +0,0 @@
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 DELETED
@@ -1,43 +0,0 @@
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 DELETED
@@ -1,1468 +0,0 @@
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 DELETED
@@ -1,159 +0,0 @@
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__)